use serde::{Deserialize, Serialize};
use crate::error::{Result, SubtitleToolkitError};
use super::{
TranslationLimits, TranslationRequest, Translator, provider_http_error,
provider_transport_error,
};
#[derive(Debug, Clone)]
pub struct OllamaTranslator {
client: reqwest::Client,
base_url: String,
model: String,
}
impl OllamaTranslator {
pub fn new(model: impl Into<String>) -> Result<Self> {
Self::with_base_url("http://localhost:11434", model)
}
pub fn with_base_url(base_url: impl Into<String>, model: impl Into<String>) -> Result<Self> {
let client = reqwest::Client::builder()
.timeout(std::time::Duration::from_secs(120))
.build()
.map_err(SubtitleToolkitError::Http)?;
Ok(Self {
client,
base_url: base_url.into().trim_end_matches('/').to_string(),
model: model.into(),
})
}
}
#[async_trait::async_trait]
impl Translator for OllamaTranslator {
async fn translate(&self, request: TranslationRequest<'_>) -> Result<String> {
let response = self
.client
.post(format!("{}/api/generate", self.base_url))
.json(&OllamaGenerateRequest {
model: &self.model,
prompt: &build_prompt(&request),
stream: false,
})
.send()
.await
.map_err(|error| provider_transport_error("ollama", error))?;
if !response.status().is_success() {
return Err(provider_http_error("ollama", response).await);
}
let body = response.json::<OllamaGenerateResponse>().await?;
if body.done != Some(true) {
return Err(SubtitleToolkitError::InvalidTranslation {
message: "ollama returned an incomplete non-streaming response".into(),
});
}
Ok(body.response.trim().to_string())
}
fn identifier(&self) -> String {
format!("ollama:{}:{}", self.base_url, self.model)
}
fn limits(&self) -> TranslationLimits {
TranslationLimits {
reject_unchanged_output: true,
..TranslationLimits::default()
}
}
}
fn build_prompt(request: &TranslationRequest<'_>) -> String {
format!(
"Translate the following subtitle dialogue to {target_language}.\n\n\
Rules:\n\
- Preserve every numeric tag exactly, like <1>, <2>, <3>.\n\
- Preserve tokens such as [[PSY_TAG_0]] exactly and in order.\n\
- Return only translated subtitle lines.\n\
- Do not add explanations, markdown, notes, or code fences.\n\
- Keep line breaks inside each subtitle when needed.\n\
- Do not add curly-brace commands or backslash formatting.\n\n\
Subtitle dialogue:\n\
{source_text}",
target_language = request.target_language,
source_text = request.source_text,
)
}
#[derive(Debug, Serialize)]
struct OllamaGenerateRequest<'a> {
model: &'a str,
prompt: &'a str,
stream: bool,
}
#[derive(Debug, Deserialize)]
struct OllamaGenerateResponse {
response: String,
#[serde(default)]
done: Option<bool>,
}
#[cfg(test)]
mod tests {
use super::*;
use wiremock::matchers::{method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
#[tokio::test]
async fn translates_numbered_text() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/api/generate"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"response": "<1> Olá\n<2> mundo",
"done": true
})))
.mount(&server)
.await;
let translator = OllamaTranslator::with_base_url(server.uri(), "test-model").unwrap();
let result = translator
.translate(TranslationRequest {
source_text: "<1> hello\n<2> world",
target_language: "pt-BR",
})
.await
.unwrap();
assert_eq!(result, "<1> Olá\n<2> mundo");
}
#[tokio::test]
async fn trims_whitespace_from_response() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"response": " <1> Olá \n",
"done": true
})))
.mount(&server)
.await;
let translator = OllamaTranslator::with_base_url(server.uri(), "test-model").unwrap();
let result = translator
.translate(TranslationRequest {
source_text: "<1> hello",
target_language: "pt-BR",
})
.await
.unwrap();
assert_eq!(result, "<1> Olá");
}
#[tokio::test]
async fn error_on_non_200() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.respond_with(ResponseTemplate::new(500).set_body_string("model not found"))
.mount(&server)
.await;
let translator = OllamaTranslator::with_base_url(server.uri(), "bad-model").unwrap();
let err = translator
.translate(TranslationRequest {
source_text: "<1> hello",
target_language: "pt-BR",
})
.await
.unwrap_err();
assert!(err.to_string().contains("ollama"));
assert!(err.to_string().contains("model not found"));
}
#[tokio::test]
async fn sends_correct_model_and_prompt() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/api/generate"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"response": "<1> ok",
"done": true
})))
.expect(1)
.mount(&server)
.await;
let translator = OllamaTranslator::with_base_url(server.uri(), "my-model").unwrap();
translator
.translate(TranslationRequest {
source_text: "<1> test",
target_language: "ja",
})
.await
.unwrap();
}
#[tokio::test]
async fn rejects_incomplete_non_streaming_response() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"response": "<1> partial",
"done": false
})))
.mount(&server)
.await;
let translator = OllamaTranslator::with_base_url(server.uri(), "model").unwrap();
let error = translator
.translate(TranslationRequest {
source_text: "<1> hello",
target_language: "pt-BR",
})
.await
.unwrap_err();
assert!(error.to_string().contains("incomplete"));
}
}