@@ -35,5 +35,12 @@ zip.workspace = true
|
|||||||
tokio = { workspace = true, features = ["macros", "rt-multi-thread"] }
|
tokio = { workspace = true, features = ["macros", "rt-multi-thread"] }
|
||||||
wiremock.workspace = true
|
wiremock.workspace = true
|
||||||
|
|
||||||
|
# Only built with the feature it exercises — cargo skips the target
|
||||||
|
# otherwise, rather than compiling an empty one that flags every crate
|
||||||
|
# dependency as unused.
|
||||||
|
[[test]]
|
||||||
|
name = "openai"
|
||||||
|
required-features = ["translate-openai"]
|
||||||
|
|
||||||
[lints]
|
[lints]
|
||||||
workspace = true
|
workspace = true
|
||||||
|
|||||||
@@ -30,6 +30,8 @@ use wiremock as _;
|
|||||||
|
|
||||||
pub mod error;
|
pub mod error;
|
||||||
pub mod model;
|
pub mod model;
|
||||||
|
#[cfg(feature = "translate-openai")]
|
||||||
|
pub mod openai;
|
||||||
pub mod opensubtitles;
|
pub mod opensubtitles;
|
||||||
pub mod podnapisi;
|
pub mod podnapisi;
|
||||||
pub mod srt;
|
pub mod srt;
|
||||||
@@ -40,6 +42,8 @@ pub use error::{Error, Result};
|
|||||||
pub use model::{
|
pub use model::{
|
||||||
Candidate, CandidateId, Fetched, MediaFile, MediaRef, ProviderId, SearchRequest, SubtitleFormat,
|
Candidate, CandidateId, Fetched, MediaFile, MediaRef, ProviderId, SearchRequest, SubtitleFormat,
|
||||||
};
|
};
|
||||||
|
#[cfg(feature = "translate-openai")]
|
||||||
|
pub use openai::{OpenAi, OpenAiConfig};
|
||||||
pub use opensubtitles::{moviehash, OpenSubtitles, OpenSubtitlesConfig};
|
pub use opensubtitles::{moviehash, OpenSubtitles, OpenSubtitlesConfig};
|
||||||
pub use podnapisi::{Podnapisi, PodnapisiBuilder, DEFAULT_BASE_URL as PODNAPISI_DEFAULT_BASE_URL};
|
pub use podnapisi::{Podnapisi, PodnapisiBuilder, DEFAULT_BASE_URL as PODNAPISI_DEFAULT_BASE_URL};
|
||||||
pub use srt::Cue;
|
pub use srt::Cue;
|
||||||
|
|||||||
@@ -0,0 +1,444 @@
|
|||||||
|
//! The OpenAI-compatible translation backend (`DESIGN.md` §15, issue #191).
|
||||||
|
//!
|
||||||
|
//! One HTTP shape — the `/chat/completions` endpoint `OpenAI` defined — reaches
|
||||||
|
//! the largest number of options: `OpenAI` itself, `OpenRouter`, Groq, a
|
||||||
|
//! self-hosted gateway, or a local `llama.cpp` server. The model and base URL
|
||||||
|
//! are constructor arguments rather than fields on [`OpenAiConfig`]: which
|
||||||
|
//! engine runs, and where, is a database setting (DESIGN.md §15, issue #198).
|
||||||
|
//! Only the credential is bootstrap config or environment, per §10 — and even
|
||||||
|
//! that is optional, since a self-hosted gateway may not ask for one.
|
||||||
|
//!
|
||||||
|
//! The batch travels as a JSON array of `{number, text}` objects, and the
|
||||||
|
//! reply is asked for in the same shape: JSON round-trips a cue's embedded
|
||||||
|
//! newlines cleanly, where a line-oriented reply format could not tell a
|
||||||
|
//! multi-line cue from two cues. [`translate::translate`] chunks, validates
|
||||||
|
//! and reassembles; this module only builds one request and turns one reply
|
||||||
|
//! into [`TranslatedCue`]s. A refusal, a content-filtered reply or a reply cut
|
||||||
|
//! short by the token limit each fail with their own [`Error::Malformed`]
|
||||||
|
//! detail rather than reaching validation as a silently short batch.
|
||||||
|
|
||||||
|
use std::fmt;
|
||||||
|
use std::time::Duration;
|
||||||
|
|
||||||
|
use reqwest::{Client, StatusCode, Url};
|
||||||
|
use serde::{Deserialize, Serialize};
|
||||||
|
|
||||||
|
use arr_core::Language;
|
||||||
|
|
||||||
|
use crate::translate::{Backend, BackendId, Batch, Error, Result, TranslateFuture, TranslatedCue};
|
||||||
|
|
||||||
|
const DEFAULT_BASE_URL: &str = "https://api.openai.com/v1/";
|
||||||
|
const BACKEND_NAME: &str = "openai";
|
||||||
|
const USER_AGENT: &str = "arr v0.1.0";
|
||||||
|
const MAX_ERROR_BODY: usize = 512;
|
||||||
|
|
||||||
|
/// Credentials for an OpenAI-compatible endpoint (DESIGN.md §10).
|
||||||
|
///
|
||||||
|
/// Never persisted to the database — it reaches here from bootstrap config or
|
||||||
|
/// environment only. Optional because a self-hosted gateway (a local
|
||||||
|
/// `llama.cpp` server, say) may not require one.
|
||||||
|
#[derive(Clone, PartialEq, Eq)]
|
||||||
|
pub struct OpenAiConfig {
|
||||||
|
/// Sent as an `Authorization: Bearer` header when present.
|
||||||
|
pub api_key: Option<String>,
|
||||||
|
}
|
||||||
|
|
||||||
|
// Hand-written: a derived Debug would print the key.
|
||||||
|
impl fmt::Debug for OpenAiConfig {
|
||||||
|
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||||
|
f.debug_struct("OpenAiConfig")
|
||||||
|
.field("api_key", &self.api_key.as_ref().map(|_| "<redacted>"))
|
||||||
|
.finish()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// The OpenAI-compatible backend behind the [`Backend`] trait.
|
||||||
|
pub struct OpenAi {
|
||||||
|
http: Client,
|
||||||
|
base_url: Url,
|
||||||
|
config: OpenAiConfig,
|
||||||
|
model: String,
|
||||||
|
id: BackendId,
|
||||||
|
}
|
||||||
|
|
||||||
|
// Hand-written: the config must never reach a log line.
|
||||||
|
impl fmt::Debug for OpenAi {
|
||||||
|
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||||
|
f.debug_struct("OpenAi")
|
||||||
|
.field("base_url", &self.base_url.as_str())
|
||||||
|
.field("model", &self.model)
|
||||||
|
.finish_non_exhaustive()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
impl OpenAi {
|
||||||
|
/// A client against `https://api.openai.com/v1/`.
|
||||||
|
///
|
||||||
|
/// # Errors
|
||||||
|
///
|
||||||
|
/// Fails if the HTTP client cannot be constructed.
|
||||||
|
pub fn new(model: impl Into<String>, config: OpenAiConfig) -> Result<Self> {
|
||||||
|
Self::with_base_url(model, config, DEFAULT_BASE_URL)
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Point the client at another OpenAI-compatible endpoint — `OpenRouter`,
|
||||||
|
/// Groq, a self-hosted gateway, a local `llama.cpp` server, or `wiremock`
|
||||||
|
/// in tests.
|
||||||
|
///
|
||||||
|
/// # Errors
|
||||||
|
///
|
||||||
|
/// Same as [`Self::new`], plus a `base_url` that does not parse.
|
||||||
|
pub fn with_base_url(
|
||||||
|
model: impl Into<String>,
|
||||||
|
config: OpenAiConfig,
|
||||||
|
base_url: &str,
|
||||||
|
) -> Result<Self> {
|
||||||
|
let id = BackendId::new(BACKEND_NAME);
|
||||||
|
let http = Client::builder().build().map_err(|err| Error::Transport {
|
||||||
|
backend: id.clone(),
|
||||||
|
source: Box::new(err),
|
||||||
|
})?;
|
||||||
|
let base_url = base_url.parse().map_err(|err| Error::Malformed {
|
||||||
|
backend: id.clone(),
|
||||||
|
detail: format!("bad base URL: {err}"),
|
||||||
|
})?;
|
||||||
|
Ok(Self {
|
||||||
|
http,
|
||||||
|
base_url,
|
||||||
|
config,
|
||||||
|
model: model.into(),
|
||||||
|
id,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
async fn check_status(&self, response: reqwest::Response) -> Result<reqwest::Response> {
|
||||||
|
match response.status() {
|
||||||
|
StatusCode::OK => Ok(response),
|
||||||
|
StatusCode::UNAUTHORIZED | StatusCode::FORBIDDEN => Err(Error::Unauthorized {
|
||||||
|
backend: self.id.clone(),
|
||||||
|
}),
|
||||||
|
StatusCode::TOO_MANY_REQUESTS => Err(Error::RateLimited {
|
||||||
|
backend: self.id.clone(),
|
||||||
|
retry_after: retry_after(&response),
|
||||||
|
}),
|
||||||
|
other => {
|
||||||
|
let body = response.text().await.unwrap_or_default();
|
||||||
|
Err(Error::Malformed {
|
||||||
|
backend: self.id.clone(),
|
||||||
|
detail: format!("status {}: {}", other.as_u16(), truncate(&body)),
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
async fn decode<T: for<'de> Deserialize<'de>>(&self, response: reqwest::Response) -> Result<T> {
|
||||||
|
let bytes = response.bytes().await.map_err(|err| Error::Transport {
|
||||||
|
backend: self.id.clone(),
|
||||||
|
source: Box::new(err),
|
||||||
|
})?;
|
||||||
|
serde_json::from_slice(&bytes).map_err(|err| Error::Malformed {
|
||||||
|
backend: self.id.clone(),
|
||||||
|
detail: err.to_string(),
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
async fn translate_inner(&self, batch: &Batch) -> Result<Vec<TranslatedCue>> {
|
||||||
|
let request_cues: Vec<RequestCue<'_>> = batch
|
||||||
|
.cues
|
||||||
|
.iter()
|
||||||
|
.map(|cue| RequestCue {
|
||||||
|
number: cue.number,
|
||||||
|
text: &cue.text,
|
||||||
|
})
|
||||||
|
.collect();
|
||||||
|
let user_content =
|
||||||
|
serde_json::to_string(&request_cues).map_err(|err| Error::Malformed {
|
||||||
|
backend: self.id.clone(),
|
||||||
|
detail: format!("could not encode batch: {err}"),
|
||||||
|
})?;
|
||||||
|
|
||||||
|
let body = ChatRequest {
|
||||||
|
model: &self.model,
|
||||||
|
messages: vec![
|
||||||
|
ChatMessage {
|
||||||
|
role: "system",
|
||||||
|
content: system_prompt(&batch.source, &batch.target),
|
||||||
|
},
|
||||||
|
ChatMessage {
|
||||||
|
role: "user",
|
||||||
|
content: user_content,
|
||||||
|
},
|
||||||
|
],
|
||||||
|
};
|
||||||
|
|
||||||
|
let url = self
|
||||||
|
.base_url
|
||||||
|
.join("chat/completions")
|
||||||
|
.map_err(|err| Error::Malformed {
|
||||||
|
backend: self.id.clone(),
|
||||||
|
detail: format!("bad chat completions path: {err}"),
|
||||||
|
})?;
|
||||||
|
|
||||||
|
let mut request = self
|
||||||
|
.http
|
||||||
|
.post(url)
|
||||||
|
.header("User-Agent", USER_AGENT)
|
||||||
|
.json(&body);
|
||||||
|
if let Some(api_key) = self.config.api_key.as_deref() {
|
||||||
|
request = request.bearer_auth(api_key);
|
||||||
|
}
|
||||||
|
|
||||||
|
tracing::debug!(backend = %self.id, model = %self.model, "OpenAI translate request");
|
||||||
|
let response = request.send().await.map_err(|err| Error::Transport {
|
||||||
|
backend: self.id.clone(),
|
||||||
|
source: Box::new(err),
|
||||||
|
})?;
|
||||||
|
let response = self.check_status(response).await?;
|
||||||
|
let reply: ChatResponse = self.decode(response).await?;
|
||||||
|
self.parse_reply(reply)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn parse_reply(&self, reply: ChatResponse) -> Result<Vec<TranslatedCue>> {
|
||||||
|
let malformed = |detail: String| Error::Malformed {
|
||||||
|
backend: self.id.clone(),
|
||||||
|
detail,
|
||||||
|
};
|
||||||
|
|
||||||
|
let choice = reply
|
||||||
|
.choices
|
||||||
|
.into_iter()
|
||||||
|
.next()
|
||||||
|
.ok_or_else(|| malformed("no choices in reply".to_owned()))?;
|
||||||
|
|
||||||
|
if let Some(refusal) = choice.message.refusal.filter(|refusal| !refusal.is_empty()) {
|
||||||
|
return Err(malformed(format!("backend refused: {refusal}")));
|
||||||
|
}
|
||||||
|
match choice.finish_reason.as_deref() {
|
||||||
|
Some("content_filter") => {
|
||||||
|
return Err(malformed("reply blocked by content filter".to_owned()))
|
||||||
|
}
|
||||||
|
Some("length") => {
|
||||||
|
return Err(malformed(
|
||||||
|
"reply was truncated before it finished".to_owned(),
|
||||||
|
))
|
||||||
|
}
|
||||||
|
_ => {}
|
||||||
|
}
|
||||||
|
|
||||||
|
let content = choice
|
||||||
|
.message
|
||||||
|
.content
|
||||||
|
.filter(|content| !content.trim().is_empty())
|
||||||
|
.ok_or_else(|| malformed("empty reply".to_owned()))?;
|
||||||
|
|
||||||
|
let cues: Vec<ReplyCue> = serde_json::from_str(strip_code_fence(&content))
|
||||||
|
.map_err(|err| malformed(format!("reply was not the expected JSON: {err}")))?;
|
||||||
|
|
||||||
|
Ok(cues
|
||||||
|
.into_iter()
|
||||||
|
.map(|cue| TranslatedCue {
|
||||||
|
number: cue.number,
|
||||||
|
text: cue.text,
|
||||||
|
})
|
||||||
|
.collect())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
impl Backend for OpenAi {
|
||||||
|
fn id(&self) -> BackendId {
|
||||||
|
self.id.clone()
|
||||||
|
}
|
||||||
|
|
||||||
|
/// True for every target except [`Language::PortugueseUnverified`]: an
|
||||||
|
/// unverified variant names no exact target to translate into, and an LLM
|
||||||
|
/// asked to hit one anyway would have to guess (DESIGN.md §15).
|
||||||
|
fn supports(&self, target: &Language) -> bool {
|
||||||
|
!matches!(target, Language::PortugueseUnverified)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn translate<'a>(&'a self, batch: &'a Batch) -> TranslateFuture<'a> {
|
||||||
|
Box::pin(async move { self.translate_inner(batch).await })
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fn retry_after(response: &reqwest::Response) -> Option<Duration> {
|
||||||
|
response
|
||||||
|
.headers()
|
||||||
|
.get(reqwest::header::RETRY_AFTER)?
|
||||||
|
.to_str()
|
||||||
|
.ok()?
|
||||||
|
.trim()
|
||||||
|
.parse::<u64>()
|
||||||
|
.ok()
|
||||||
|
.map(Duration::from_secs)
|
||||||
|
}
|
||||||
|
|
||||||
|
fn truncate(body: &str) -> String {
|
||||||
|
body.chars().take(MAX_ERROR_BODY).collect()
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Strip a ```` ```json ```` / ```` ``` ```` fence a model wrapped its reply
|
||||||
|
/// in despite being told not to. Cheaper than retrying the whole batch over a
|
||||||
|
/// habit some models never break.
|
||||||
|
fn strip_code_fence(content: &str) -> &str {
|
||||||
|
let trimmed = content.trim();
|
||||||
|
let trimmed = trimmed
|
||||||
|
.strip_prefix("```json")
|
||||||
|
.or_else(|| trimmed.strip_prefix("```"))
|
||||||
|
.unwrap_or(trimmed);
|
||||||
|
trimmed.strip_suffix("```").unwrap_or(trimmed).trim()
|
||||||
|
}
|
||||||
|
|
||||||
|
/// How a language reads in a prompt meant for an LLM, not a settings row.
|
||||||
|
fn language_name(language: &Language) -> String {
|
||||||
|
match language {
|
||||||
|
Language::PortuguesePortugal => "European Portuguese (pt-PT)".to_owned(),
|
||||||
|
Language::PortugueseBrazil => "Brazilian Portuguese (pt-BR)".to_owned(),
|
||||||
|
Language::PortugueseUnverified => "Portuguese".to_owned(),
|
||||||
|
Language::Other(tag) => tag.clone(),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
/// A concrete nudge against the one mistake DESIGN.md §15 calls out by name:
|
||||||
|
/// answering the wrong Portuguese. Generic advice ("respect the variant")
|
||||||
|
/// is easy for a model to nod at and ignore; naming the words it must not
|
||||||
|
/// reach for is not.
|
||||||
|
fn variant_note(target: &Language) -> Option<&'static str> {
|
||||||
|
match target {
|
||||||
|
Language::PortuguesePortugal => Some(
|
||||||
|
"This is continental European Portuguese, not Brazilian Portuguese: write \"tu\", \
|
||||||
|
\"comboio\", \"casa de banho\" and European spelling and verb forms, never \
|
||||||
|
Brazilian ones.",
|
||||||
|
),
|
||||||
|
Language::PortugueseBrazil => Some(
|
||||||
|
"This is Brazilian Portuguese, not European Portuguese: write \"você\", \"trem\", \
|
||||||
|
\"banheiro\" and Brazilian spelling and verb forms, never European ones.",
|
||||||
|
),
|
||||||
|
_ => None,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fn system_prompt(source: &Language, target: &Language) -> String {
|
||||||
|
let mut prompt = format!(
|
||||||
|
"You translate subtitle cues for a film or TV episode from {} into {}. Write natural, \
|
||||||
|
spoken subtitle language a native speaker would actually use, not a stiff literal \
|
||||||
|
translation.",
|
||||||
|
language_name(source),
|
||||||
|
language_name(target)
|
||||||
|
);
|
||||||
|
if let Some(note) = variant_note(target) {
|
||||||
|
prompt.push(' ');
|
||||||
|
prompt.push_str(note);
|
||||||
|
}
|
||||||
|
prompt.push_str(
|
||||||
|
" Reply with nothing but a JSON array, one object per cue, each shaped \
|
||||||
|
{\"number\": <the same number you were given>, \"text\": \"<translation>\"}, in the \
|
||||||
|
order you received them. Every cue you were given must appear exactly once, under its \
|
||||||
|
original number: never add, drop, merge or renumber a cue. No prose, no markdown code \
|
||||||
|
fences, no explanation — JSON only.",
|
||||||
|
);
|
||||||
|
prompt
|
||||||
|
}
|
||||||
|
|
||||||
|
#[derive(Serialize)]
|
||||||
|
struct RequestCue<'a> {
|
||||||
|
number: usize,
|
||||||
|
text: &'a str,
|
||||||
|
}
|
||||||
|
|
||||||
|
#[derive(Deserialize)]
|
||||||
|
struct ReplyCue {
|
||||||
|
number: usize,
|
||||||
|
text: String,
|
||||||
|
}
|
||||||
|
|
||||||
|
#[derive(Serialize)]
|
||||||
|
struct ChatMessage {
|
||||||
|
role: &'static str,
|
||||||
|
content: String,
|
||||||
|
}
|
||||||
|
|
||||||
|
#[derive(Serialize)]
|
||||||
|
struct ChatRequest<'a> {
|
||||||
|
model: &'a str,
|
||||||
|
messages: Vec<ChatMessage>,
|
||||||
|
}
|
||||||
|
|
||||||
|
#[derive(Debug, Deserialize)]
|
||||||
|
struct ChatResponse {
|
||||||
|
choices: Vec<ChatChoice>,
|
||||||
|
}
|
||||||
|
|
||||||
|
#[derive(Debug, Deserialize)]
|
||||||
|
struct ChatChoice {
|
||||||
|
message: ChatReplyMessage,
|
||||||
|
finish_reason: Option<String>,
|
||||||
|
}
|
||||||
|
|
||||||
|
#[derive(Debug, Deserialize)]
|
||||||
|
struct ChatReplyMessage {
|
||||||
|
content: Option<String>,
|
||||||
|
#[serde(default)]
|
||||||
|
refusal: Option<String>,
|
||||||
|
}
|
||||||
|
|
||||||
|
#[cfg(test)]
|
||||||
|
mod tests {
|
||||||
|
use arr_core::Language;
|
||||||
|
|
||||||
|
use super::{language_name, strip_code_fence, system_prompt, OpenAi};
|
||||||
|
use crate::translate::Backend;
|
||||||
|
|
||||||
|
fn backend() -> OpenAi {
|
||||||
|
OpenAi::new("gpt-4o-mini", super::OpenAiConfig { api_key: None }).expect("client builds")
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn every_target_but_unverified_portuguese_is_supported() {
|
||||||
|
let backend = backend();
|
||||||
|
assert!(backend.supports(&Language::PortuguesePortugal));
|
||||||
|
assert!(backend.supports(&Language::PortugueseBrazil));
|
||||||
|
assert!(backend.supports(&Language::Other("en".to_owned())));
|
||||||
|
assert!(!backend.supports(&Language::PortugueseUnverified));
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn the_prompt_tells_pt_pt_and_pt_br_apart() {
|
||||||
|
let source = Language::Other("en".to_owned());
|
||||||
|
|
||||||
|
let pt_pt = system_prompt(&source, &Language::PortuguesePortugal);
|
||||||
|
let pt_br = system_prompt(&source, &Language::PortugueseBrazil);
|
||||||
|
|
||||||
|
assert!(pt_pt.contains("European Portuguese"));
|
||||||
|
assert!(pt_pt.contains("comboio"));
|
||||||
|
assert!(pt_br.contains("Brazilian Portuguese"));
|
||||||
|
assert!(pt_br.contains("trem"));
|
||||||
|
assert_ne!(pt_pt, pt_br);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn the_prompt_asks_for_json_only() {
|
||||||
|
let prompt = system_prompt(
|
||||||
|
&Language::Other("en".to_owned()),
|
||||||
|
&Language::Other("fr".to_owned()),
|
||||||
|
);
|
||||||
|
assert!(prompt.contains("JSON"));
|
||||||
|
assert!(prompt.contains("never add, drop, merge or renumber"));
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn language_names_read_like_prose_not_settings_rows() {
|
||||||
|
assert_eq!(
|
||||||
|
language_name(&Language::PortuguesePortugal),
|
||||||
|
"European Portuguese (pt-PT)"
|
||||||
|
);
|
||||||
|
assert_eq!(language_name(&Language::Other("en".to_owned())), "en");
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn a_fenced_reply_is_unwrapped() {
|
||||||
|
assert_eq!(strip_code_fence("```json\n[1,2]\n```"), "[1,2]");
|
||||||
|
assert_eq!(strip_code_fence("```\n[1,2]\n```"), "[1,2]");
|
||||||
|
assert_eq!(strip_code_fence("[1,2]"), "[1,2]");
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,328 @@
|
|||||||
|
//! OpenAI-compatible backend tests (`DESIGN.md` §15, issue #191). Everything
|
||||||
|
//! runs against `wiremock` — the live service is never touched.
|
||||||
|
//!
|
||||||
|
//! Built only with the `translate-openai` feature (`required-features` in
|
||||||
|
//! `Cargo.toml`), same as the module itself: with the feature off, cargo
|
||||||
|
//! skips this target rather than building an empty one.
|
||||||
|
|
||||||
|
// Same per-target quirk as the crate's own tests: an integration test links
|
||||||
|
// the library's dependencies without using them all directly.
|
||||||
|
use {
|
||||||
|
arr_parse as _, chardetng as _, encoding_rs as _, reqwest as _, serde as _, serde_json as _,
|
||||||
|
tempfile as _, thiserror as _, tracing as _, zip as _,
|
||||||
|
};
|
||||||
|
|
||||||
|
use std::time::Duration;
|
||||||
|
|
||||||
|
use arr_core::Language;
|
||||||
|
use arr_subs::translate::{translate, Backend, Error};
|
||||||
|
use arr_subs::{OpenAi, OpenAiConfig};
|
||||||
|
use wiremock::matchers::{header, method, path};
|
||||||
|
use wiremock::{Mock, MockServer, ResponseTemplate};
|
||||||
|
|
||||||
|
fn backend(server: &MockServer, api_key: Option<&str>) -> OpenAi {
|
||||||
|
OpenAi::with_base_url(
|
||||||
|
"gpt-4o-mini",
|
||||||
|
OpenAiConfig {
|
||||||
|
api_key: api_key.map(str::to_owned),
|
||||||
|
},
|
||||||
|
&format!("{}/v1/", server.uri()),
|
||||||
|
)
|
||||||
|
.expect("client builds")
|
||||||
|
}
|
||||||
|
|
||||||
|
fn chat_response(body: &str) -> String {
|
||||||
|
format!(r#"{{"choices": [{{"message": {{"content": {body}}}, "finish_reason": "stop"}}]}}"#)
|
||||||
|
}
|
||||||
|
|
||||||
|
async fn mount_reply(server: &MockServer, response: ResponseTemplate) {
|
||||||
|
Mock::given(method("POST"))
|
||||||
|
.and(path("/v1/chat/completions"))
|
||||||
|
.respond_with(response)
|
||||||
|
.mount(server)
|
||||||
|
.await;
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn a_batch_translates_and_timings_never_move() {
|
||||||
|
let server = MockServer::start().await;
|
||||||
|
let body = chat_response(
|
||||||
|
r#""[{\"number\":1,\"text\":\"ol\\u00e1\"},{\"number\":2,\"text\":\"mundo\"}]""#,
|
||||||
|
);
|
||||||
|
mount_reply(&server, ResponseTemplate::new(200).set_body_string(body)).await;
|
||||||
|
let backend = backend(&server, Some("sk-test"));
|
||||||
|
|
||||||
|
let cues = vec![
|
||||||
|
arr_subs::Cue {
|
||||||
|
start: Duration::from_secs(0),
|
||||||
|
end: Duration::from_secs(1),
|
||||||
|
text: "hello".to_owned(),
|
||||||
|
},
|
||||||
|
arr_subs::Cue {
|
||||||
|
start: Duration::from_secs(1),
|
||||||
|
end: Duration::from_secs(2),
|
||||||
|
text: "world".to_owned(),
|
||||||
|
},
|
||||||
|
];
|
||||||
|
|
||||||
|
let translated = translate(
|
||||||
|
&backend,
|
||||||
|
&cues,
|
||||||
|
&Language::Other("en".to_owned()),
|
||||||
|
&Language::PortuguesePortugal,
|
||||||
|
)
|
||||||
|
.await
|
||||||
|
.expect("the mocked reply is well-formed");
|
||||||
|
|
||||||
|
assert_eq!(translated[0].text, "ol\u{e1}");
|
||||||
|
assert_eq!(translated[1].text, "mundo");
|
||||||
|
assert_eq!(translated[0].start, cues[0].start);
|
||||||
|
assert_eq!(translated[1].end, cues[1].end);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn the_api_key_travels_as_a_bearer_header() {
|
||||||
|
let server = MockServer::start().await;
|
||||||
|
Mock::given(method("POST"))
|
||||||
|
.and(path("/v1/chat/completions"))
|
||||||
|
.and(header("Authorization", "Bearer sk-test"))
|
||||||
|
.respond_with(
|
||||||
|
ResponseTemplate::new(200)
|
||||||
|
.set_body_string(chat_response(r#""[{\"number\":1,\"text\":\"oi\"}]""#)),
|
||||||
|
)
|
||||||
|
.mount(&server)
|
||||||
|
.await;
|
||||||
|
let backend = backend(&server, Some("sk-test"));
|
||||||
|
|
||||||
|
let cues = vec![arr_subs::Cue {
|
||||||
|
start: Duration::from_secs(0),
|
||||||
|
end: Duration::from_secs(1),
|
||||||
|
text: "hi".to_owned(),
|
||||||
|
}];
|
||||||
|
|
||||||
|
translate(
|
||||||
|
&backend,
|
||||||
|
&cues,
|
||||||
|
&Language::Other("en".to_owned()),
|
||||||
|
&Language::PortuguesePortugal,
|
||||||
|
)
|
||||||
|
.await
|
||||||
|
.expect("the mock only answers a bearer-authed request");
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn a_fenced_json_reply_still_parses() {
|
||||||
|
let server = MockServer::start().await;
|
||||||
|
let body = chat_response(r#""```json\n[{\"number\":1,\"text\":\"oi\"}]\n```""#);
|
||||||
|
mount_reply(&server, ResponseTemplate::new(200).set_body_string(body)).await;
|
||||||
|
let backend = backend(&server, None);
|
||||||
|
|
||||||
|
let cues = vec![arr_subs::Cue {
|
||||||
|
start: Duration::from_secs(0),
|
||||||
|
end: Duration::from_secs(1),
|
||||||
|
text: "hi".to_owned(),
|
||||||
|
}];
|
||||||
|
|
||||||
|
let translated = translate(
|
||||||
|
&backend,
|
||||||
|
&cues,
|
||||||
|
&Language::Other("en".to_owned()),
|
||||||
|
&Language::PortuguesePortugal,
|
||||||
|
)
|
||||||
|
.await
|
||||||
|
.expect("a markdown fence around the JSON is still accepted");
|
||||||
|
|
||||||
|
assert_eq!(translated[0].text, "oi");
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn unauthorized_status_is_reported_as_such() {
|
||||||
|
let server = MockServer::start().await;
|
||||||
|
mount_reply(&server, ResponseTemplate::new(401)).await;
|
||||||
|
let backend = backend(&server, Some("sk-bad"));
|
||||||
|
|
||||||
|
let error = translate(
|
||||||
|
&backend,
|
||||||
|
&[arr_subs::Cue {
|
||||||
|
start: Duration::from_secs(0),
|
||||||
|
end: Duration::from_secs(1),
|
||||||
|
text: "hi".to_owned(),
|
||||||
|
}],
|
||||||
|
&Language::Other("en".to_owned()),
|
||||||
|
&Language::PortuguesePortugal,
|
||||||
|
)
|
||||||
|
.await
|
||||||
|
.expect_err("401 means the key was rejected");
|
||||||
|
|
||||||
|
assert!(matches!(error, Error::Unauthorized { .. }));
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn a_rate_limit_carries_its_retry_after() {
|
||||||
|
let server = MockServer::start().await;
|
||||||
|
mount_reply(
|
||||||
|
&server,
|
||||||
|
ResponseTemplate::new(429).insert_header("retry-after", "30"),
|
||||||
|
)
|
||||||
|
.await;
|
||||||
|
let backend = backend(&server, None);
|
||||||
|
|
||||||
|
let error = translate(
|
||||||
|
&backend,
|
||||||
|
&[arr_subs::Cue {
|
||||||
|
start: Duration::from_secs(0),
|
||||||
|
end: Duration::from_secs(1),
|
||||||
|
text: "hi".to_owned(),
|
||||||
|
}],
|
||||||
|
&Language::Other("en".to_owned()),
|
||||||
|
&Language::PortuguesePortugal,
|
||||||
|
)
|
||||||
|
.await
|
||||||
|
.expect_err("429 is a rate limit, not a broken backend");
|
||||||
|
|
||||||
|
match error {
|
||||||
|
Error::RateLimited { retry_after, .. } => {
|
||||||
|
assert_eq!(retry_after, Some(Duration::from_secs(30)));
|
||||||
|
}
|
||||||
|
other => panic!("expected RateLimited, got {other:?}"),
|
||||||
|
}
|
||||||
|
assert!(error_is_transient(&error));
|
||||||
|
}
|
||||||
|
|
||||||
|
fn error_is_transient(error: &Error) -> bool {
|
||||||
|
error.is_transient()
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn a_content_filtered_reply_is_an_error_not_a_silent_partial() {
|
||||||
|
let server = MockServer::start().await;
|
||||||
|
let body =
|
||||||
|
r#"{"choices": [{"message": {"content": null}, "finish_reason": "content_filter"}]}"#;
|
||||||
|
mount_reply(&server, ResponseTemplate::new(200).set_body_string(body)).await;
|
||||||
|
let backend = backend(&server, None);
|
||||||
|
|
||||||
|
let error = translate(
|
||||||
|
&backend,
|
||||||
|
&[arr_subs::Cue {
|
||||||
|
start: Duration::from_secs(0),
|
||||||
|
end: Duration::from_secs(1),
|
||||||
|
text: "hi".to_owned(),
|
||||||
|
}],
|
||||||
|
&Language::Other("en".to_owned()),
|
||||||
|
&Language::PortuguesePortugal,
|
||||||
|
)
|
||||||
|
.await
|
||||||
|
.expect_err("a content-filtered reply must not translate to nothing silently");
|
||||||
|
|
||||||
|
assert!(matches!(error, Error::Malformed { .. }));
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn a_truncated_reply_is_an_error_not_a_silent_partial() {
|
||||||
|
let server = MockServer::start().await;
|
||||||
|
let body = chat_response(r#""[{\"number\":1,\"text\":\"oi""#);
|
||||||
|
let body = body.replace(r#""finish_reason": "stop""#, r#""finish_reason": "length""#);
|
||||||
|
mount_reply(&server, ResponseTemplate::new(200).set_body_string(body)).await;
|
||||||
|
let backend = backend(&server, None);
|
||||||
|
|
||||||
|
let error = translate(
|
||||||
|
&backend,
|
||||||
|
&[arr_subs::Cue {
|
||||||
|
start: Duration::from_secs(0),
|
||||||
|
end: Duration::from_secs(1),
|
||||||
|
text: "hi".to_owned(),
|
||||||
|
}],
|
||||||
|
&Language::Other("en".to_owned()),
|
||||||
|
&Language::PortuguesePortugal,
|
||||||
|
)
|
||||||
|
.await
|
||||||
|
.expect_err("a length-truncated reply must fail loudly");
|
||||||
|
|
||||||
|
assert!(matches!(error, Error::Malformed { .. }));
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn a_refusal_is_an_error_not_a_silent_partial() {
|
||||||
|
let server = MockServer::start().await;
|
||||||
|
let body = r#"{"choices": [{"message": {"content": null, "refusal": "I can't help with that."}, "finish_reason": "stop"}]}"#;
|
||||||
|
mount_reply(&server, ResponseTemplate::new(200).set_body_string(body)).await;
|
||||||
|
let backend = backend(&server, None);
|
||||||
|
|
||||||
|
let error = translate(
|
||||||
|
&backend,
|
||||||
|
&[arr_subs::Cue {
|
||||||
|
start: Duration::from_secs(0),
|
||||||
|
end: Duration::from_secs(1),
|
||||||
|
text: "hi".to_owned(),
|
||||||
|
}],
|
||||||
|
&Language::Other("en".to_owned()),
|
||||||
|
&Language::PortuguesePortugal,
|
||||||
|
)
|
||||||
|
.await
|
||||||
|
.expect_err("a refusal must not be read as an empty translation");
|
||||||
|
|
||||||
|
assert!(matches!(error, Error::Malformed { .. }));
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn a_mismatched_reply_is_rejected_by_shared_validation() {
|
||||||
|
// The backend does not validate its own replies (DESIGN.md §15) — this
|
||||||
|
// proves the shared `translate()` step catches a dropped cue even when
|
||||||
|
// the backend itself is this one.
|
||||||
|
let server = MockServer::start().await;
|
||||||
|
let body = chat_response(r#""[{\"number\":1,\"text\":\"oi\"}]""#);
|
||||||
|
mount_reply(&server, ResponseTemplate::new(200).set_body_string(body)).await;
|
||||||
|
let backend = backend(&server, None);
|
||||||
|
|
||||||
|
let error = translate(
|
||||||
|
&backend,
|
||||||
|
&[
|
||||||
|
arr_subs::Cue {
|
||||||
|
start: Duration::from_secs(0),
|
||||||
|
end: Duration::from_secs(1),
|
||||||
|
text: "hi".to_owned(),
|
||||||
|
},
|
||||||
|
arr_subs::Cue {
|
||||||
|
start: Duration::from_secs(1),
|
||||||
|
end: Duration::from_secs(2),
|
||||||
|
text: "there".to_owned(),
|
||||||
|
},
|
||||||
|
],
|
||||||
|
&Language::Other("en".to_owned()),
|
||||||
|
&Language::PortuguesePortugal,
|
||||||
|
)
|
||||||
|
.await
|
||||||
|
.expect_err("one cue back for two sent must be rejected");
|
||||||
|
|
||||||
|
assert!(matches!(error, Error::CueMismatch { .. }));
|
||||||
|
}
|
||||||
|
|
||||||
|
#[tokio::test]
|
||||||
|
async fn an_unsupported_target_is_refused_before_any_request() {
|
||||||
|
let server = MockServer::start().await;
|
||||||
|
// No mock mounted: a request would fail wiremock's "no matching mock".
|
||||||
|
let backend = backend(&server, None);
|
||||||
|
|
||||||
|
let error = translate(
|
||||||
|
&backend,
|
||||||
|
&[arr_subs::Cue {
|
||||||
|
start: Duration::from_secs(0),
|
||||||
|
end: Duration::from_secs(1),
|
||||||
|
text: "hi".to_owned(),
|
||||||
|
}],
|
||||||
|
&Language::Other("en".to_owned()),
|
||||||
|
&Language::PortugueseUnverified,
|
||||||
|
)
|
||||||
|
.await
|
||||||
|
.expect_err("an unverified variant names no target to hit");
|
||||||
|
|
||||||
|
assert!(matches!(error, Error::UnsupportedTarget { .. }));
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn the_backend_answers_to_openai() {
|
||||||
|
let id = OpenAi::new("gpt-4o-mini", OpenAiConfig { api_key: None })
|
||||||
|
.expect("client builds")
|
||||||
|
.id();
|
||||||
|
assert_eq!(id.as_str(), "openai");
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user