use crate::client::{GenerationHints, LLMClient, LLMResponse, ModelParams, TokenUsage};
use crate::coordinator::{ConversationMessage, MessageRole};
use ares_types::types::{AppError, Result, ToolCall, ToolDefinition};
use async_stream::stream;
use async_trait::async_trait;
use futures::{Stream, StreamExt};
use ollama_rs::{
Ollama,
generation::chat::{ChatMessage, request::ChatMessageRequest},
generation::tools::{ToolCall as OllamaToolCall, ToolFunctionInfo, ToolInfo, ToolType},
models::ModelOptions,
};
use schemars::Schema;
use std::sync::RwLock;
pub struct OllamaClient {
client: Ollama,
model: String,
params: ModelParams,
hints: RwLock<GenerationHints>,
}
impl OllamaClient {
pub async fn new(base_url: String, model: String) -> Result<Self> {
Self::with_params(base_url, model, ModelParams::default()).await
}
pub async fn with_params(base_url: String, model: String, params: ModelParams) -> Result<Self> {
let trimmed = base_url.trim();
if trimmed.is_empty() {
return Err(AppError::Configuration(
"OLLAMA_URL is empty/invalid; expected something like http://localhost:11434"
.to_string(),
));
}
let without_scheme = trimmed
.strip_prefix("http://")
.or_else(|| trimmed.strip_prefix("https://"))
.unwrap_or(trimmed);
let host_port = without_scheme
.split(&['/', '?', '#'][..])
.next()
.unwrap_or("localhost:11434");
let (host, port) = if let Some(colon_idx) = host_port.rfind(':') {
let h = &host_port[..colon_idx];
let p_str = &host_port[colon_idx + 1..];
let p = p_str.parse::<u16>().map_err(|_| {
AppError::Configuration(format!(
"Invalid OLLAMA_URL port in '{}'; expected e.g. http://localhost:11434",
base_url
))
})?;
(h.to_string(), p)
} else {
(host_port.to_string(), 11434)
};
let client = Ollama::builder()
.host(format!("http://{}", host))
.port(port)
.build();
Ok(Self {
client,
model,
params,
hints: RwLock::new(GenerationHints::default()),
})
}
fn hint_snapshot(&self) -> GenerationHints {
self.hints.read().map(|h| h.clone()).unwrap_or_default()
}
fn build_request(
&self,
messages: Vec<ChatMessage>,
caps: Option<&OllamaModelCapabilities>,
) -> ChatMessageRequest {
let hints = self.hint_snapshot();
let mut request = ChatMessageRequest::new(self.model.clone(), messages)
.options(self.build_model_options());
if let Some(budget) = hints.max_tokens {
let mut options = ModelOptions::default();
if let Some(temp) = self.params.temperature {
options = options.temperature(temp);
}
options = options.num_predict(budget as i32);
if let Some(top_p) = self.params.top_p {
options = options.top_p(top_p);
}
if let Some(pres_penalty) = self.params.presence_penalty {
options = options.repeat_penalty(pres_penalty);
}
request.options = Some(options);
}
if hints.json_mode && caps.is_some_and(|c| c.supports_json_mode) {
request.format = Some(ollama_rs::generation::parameters::FormatType::Json);
}
request
}
fn build_model_options(&self) -> ModelOptions {
let mut options = ModelOptions::default();
if let Some(temp) = self.params.temperature {
options = options.temperature(temp);
}
if let Some(max_tokens) = self.params.max_tokens {
options = options.num_predict(max_tokens as i32);
}
if let Some(top_p) = self.params.top_p {
options = options.top_p(top_p);
}
if let Some(pres_penalty) = self.params.presence_penalty {
options = options.repeat_penalty(pres_penalty);
}
options
}
fn convert_tool_definition(tool: &ToolDefinition) -> ToolInfo {
let schema: Schema =
serde_json::from_value(tool.parameters.clone()).unwrap_or_else(|_| Schema::default());
ToolInfo {
tool_type: ToolType::Function,
function: ToolFunctionInfo {
name: tool.name.clone(),
description: tool.description.clone(),
parameters: schema,
},
}
}
fn convert_tool_call(call: &OllamaToolCall) -> ToolCall {
ToolCall {
id: uuid::Uuid::new_v4().to_string(),
name: call.function.name.clone(),
arguments: coerce_tool_argument_types(&call.function.arguments),
}
}
fn map_ollama_error(err: &ollama_rs::error::OllamaError) -> AppError {
use ollama_rs::error::OllamaError;
match err {
OllamaError::InternalError(inner) => map_ollama_error_message(&inner.message),
OllamaError::Other(msg) => map_ollama_error_message(msg),
OllamaError::ReqwestError(e) => {
if e.is_timeout() {
AppError::LLM("Ollama request timed out".to_string())
} else if e.is_connect() {
AppError::LLM("Ollama connection failed".to_string())
} else {
map_ollama_error_message(&e.to_string())
}
}
OllamaError::JsonError(e) => AppError::LLM(format!("Ollama JSON error: {e}")),
OllamaError::ToolCallError(e) => {
AppError::InvalidInput(format!("Ollama tool error: {e}"))
}
}
}
fn convert_conversation_message(&self, msg: &ConversationMessage) -> ChatMessage {
match msg.role {
MessageRole::System => ChatMessage::system(msg.content.clone()),
MessageRole::User => ChatMessage::user(msg.content.clone()),
MessageRole::Assistant => {
ChatMessage::assistant(msg.content.clone())
}
MessageRole::Tool => {
ChatMessage::tool(msg.content.clone())
}
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct OllamaModelCapabilities {
pub supports_embeddings: bool,
pub context_window: Option<u32>,
pub supports_json_mode: bool,
}
#[allow(dead_code)]
pub(crate) fn map_ollama_http_status(
status: u16,
body: &str,
retry_after_secs: Option<u64>,
) -> AppError {
let retry_hint = retry_after_secs
.map(|secs| format!(" (retry after {secs}s)"))
.unwrap_or_default();
match status {
429 => AppError::RateLimited(format!("{body}{retry_hint}")),
401 | 403 => AppError::Auth(body.to_string()),
404 => AppError::NotFound(body.to_string()),
408 | 504 => AppError::LLM(format!("Ollama timeout: {body}")),
500..=599 => AppError::External(format!("Ollama server error ({status}): {body}")),
_ => map_ollama_error_message(body),
}
}
#[allow(dead_code)]
pub(crate) fn parse_retry_after(header: &str) -> Option<u64> {
let trimmed = header.trim();
if trimmed.is_empty() {
return None;
}
trimmed.parse::<u64>().ok()
}
pub(crate) fn parse_model_capabilities_from_show(
info: &serde_json::Value,
) -> OllamaModelCapabilities {
let parameters = info
.get("parameters")
.and_then(|v| v.as_str())
.unwrap_or("");
let modelfile = info.get("modelfile").and_then(|v| v.as_str()).unwrap_or("");
let template = info.get("template").and_then(|v| v.as_str()).unwrap_or("");
let combined = format!(
"{parameters}
{modelfile}
{template}"
);
let lower = combined.to_ascii_lowercase();
let supports_embeddings = lower.contains("embedding")
|| lower.contains("embed")
|| info
.get("capabilities")
.and_then(|c| c.as_array())
.map(|arr| {
arr.iter().any(|v| {
v.as_str()
.map(|s| {
let s = s.to_ascii_lowercase();
s.contains("embed")
})
.unwrap_or(false)
})
})
.unwrap_or(false);
let supports_json_mode =
lower.contains("format json") || (lower.contains("format") && lower.contains("json"));
OllamaModelCapabilities {
supports_embeddings,
context_window: parse_num_ctx_parameter(&combined),
supports_json_mode,
}
}
fn parse_num_ctx_parameter(text: &str) -> Option<u32> {
for line in text.lines() {
let line = line.trim();
if let Some(rest) = line.strip_prefix("num_ctx") {
let num_str = rest.split_whitespace().next().unwrap_or(rest.trim());
if let Ok(n) = num_str.parse::<u32>() {
return Some(n);
}
}
}
None
}
pub(crate) fn coerce_tool_argument_types(value: &serde_json::Value) -> serde_json::Value {
match value {
serde_json::Value::Object(map) => {
let mut out = serde_json::Map::new();
for (k, v) in map {
out.insert(k.clone(), coerce_tool_argument_types(v));
}
serde_json::Value::Object(out)
}
serde_json::Value::Array(arr) => {
serde_json::Value::Array(arr.iter().map(coerce_tool_argument_types).collect())
}
serde_json::Value::String(s) => {
if s == "true" || s == "false" {
return serde_json::Value::Bool(s == "true");
}
if let Ok(n) = s.parse::<i64>() {
return serde_json::json!(n);
}
if let Ok(n) = s.parse::<f64>() {
if s.contains('.') {
return serde_json::json!(n);
}
}
serde_json::Value::String(s.clone())
}
other => other.clone(),
}
}
pub(crate) fn resolve_finish_reason(has_tool_calls: bool, done: bool) -> &'static str {
if has_tool_calls {
"tool_calls"
} else if done {
"stop"
} else {
"length"
}
}
pub(crate) fn classify_stream_sse_failure(kind: &str) -> AppError {
let lower = kind.to_ascii_lowercase();
if lower.contains("timeout") {
AppError::LLM("Ollama stream timeout".to_string())
} else if lower.contains("disconnect")
|| lower.contains("reset")
|| lower.contains("broken pipe")
{
AppError::LLM("Ollama stream disconnected".to_string())
} else if lower.contains("json") || lower.contains("deserialize") || lower.contains("malformed")
{
AppError::LLM("Ollama stream malformed JSON".to_string())
} else {
AppError::LLM("Stream chunk error".to_string())
}
}
fn map_ollama_error_message(msg: &str) -> AppError {
let lower = msg.to_ascii_lowercase();
if lower.contains("rate limit") || lower.contains("too many requests") || lower.contains("429")
{
return AppError::RateLimited(msg.to_string());
}
if lower.contains("unauthorized")
|| lower.contains("401")
|| lower.contains("forbidden")
|| lower.contains("403")
{
return AppError::Auth(msg.to_string());
}
if lower.contains("not found") || lower.contains("404") {
return AppError::NotFound(msg.to_string());
}
if lower.contains("timeout") || lower.contains("timed out") {
return AppError::LLM(format!("Ollama timeout: {msg}"));
}
if lower.contains("connection reset")
|| lower.contains("disconnect")
|| lower.contains("broken pipe")
{
return AppError::LLM(format!("Ollama disconnected: {msg}"));
}
AppError::LLM(format!("Ollama error: {msg}"))
}
#[async_trait]
impl LLMClient for OllamaClient {
async fn generate(&self, prompt: &str) -> Result<String> {
let messages = vec![ChatMessage::user(prompt.to_string())];
let request = self.build_request(messages, None);
let response = self
.client
.send_chat_messages(request)
.await
.map_err(|e| Self::map_ollama_error(&e))?;
Ok(response.message.content)
}
async fn generate_with_system(&self, system: &str, prompt: &str) -> Result<String> {
let messages = vec![
ChatMessage::system(system.to_string()),
ChatMessage::user(prompt.to_string()),
];
let request = self.build_request(messages, None);
let response = self
.client
.send_chat_messages(request)
.await
.map_err(|e| Self::map_ollama_error(&e))?;
Ok(response.message.content)
}
async fn generate_with_history(&self, messages: &[(String, String)]) -> Result<LLMResponse> {
let chat_messages: Vec<ChatMessage> = messages
.iter()
.map(|(role, content)| match role.as_str() {
"system" => ChatMessage::system(content.clone()),
"user" => ChatMessage::user(content.clone()),
"assistant" => ChatMessage::assistant(content.clone()),
_ => ChatMessage::user(content.clone()),
})
.collect();
let request = self.build_request(chat_messages, None);
let response = self
.client
.send_chat_messages(request)
.await
.map_err(|e| Self::map_ollama_error(&e))?;
let usage = response
.final_data
.as_ref()
.map(|data| TokenUsage::new(data.prompt_eval_count as u32, data.eval_count as u32));
Ok(LLMResponse {
content: response.message.content,
tool_calls: vec![],
finish_reason: resolve_finish_reason(false, response.done).to_string(),
usage,
})
}
async fn generate_with_tools(
&self,
prompt: &str,
tools: &[ToolDefinition],
) -> Result<LLMResponse> {
let ollama_tools: Vec<ToolInfo> = tools.iter().map(Self::convert_tool_definition).collect();
let messages = vec![ChatMessage::user(prompt.to_string())];
let request = self.build_request(messages, None).tools(ollama_tools);
let response = self
.client
.send_chat_messages(request)
.await
.map_err(|e| Self::map_ollama_error(&e))?;
let content = response.message.content.clone();
let tool_calls: Vec<ToolCall> = response
.message
.tool_calls
.iter()
.map(Self::convert_tool_call)
.collect();
let usage = response
.final_data
.as_ref()
.map(|data| TokenUsage::new(data.prompt_eval_count as u32, data.eval_count as u32));
let finish_reason =
resolve_finish_reason(!tool_calls.is_empty(), response.done).to_string();
Ok(LLMResponse {
content,
tool_calls,
finish_reason,
usage,
})
}
async fn generate_with_tools_and_history(
&self,
messages: &[ConversationMessage],
tools: &[ToolDefinition],
) -> Result<LLMResponse> {
let ollama_tools: Vec<ToolInfo> = tools.iter().map(Self::convert_tool_definition).collect();
let chat_messages: Vec<ChatMessage> = messages
.iter()
.map(|msg| self.convert_conversation_message(msg))
.collect();
let mut request = self.build_request(chat_messages, None);
if !ollama_tools.is_empty() {
request = request.tools(ollama_tools);
}
let response = self
.client
.send_chat_messages(request)
.await
.map_err(|e| Self::map_ollama_error(&e))?;
let content = response.message.content.clone();
let tool_calls: Vec<ToolCall> = response
.message
.tool_calls
.iter()
.map(Self::convert_tool_call)
.collect();
let usage = response
.final_data
.as_ref()
.map(|data| TokenUsage::new(data.prompt_eval_count as u32, data.eval_count as u32));
let finish_reason =
resolve_finish_reason(!tool_calls.is_empty(), response.done).to_string();
Ok(LLMResponse {
content,
tool_calls,
finish_reason,
usage,
})
}
async fn stream(
&self,
prompt: &str,
) -> Result<Box<dyn Stream<Item = Result<String>> + Send + Unpin>> {
let messages = vec![ChatMessage::user(prompt.to_string())];
let request = self.build_request(messages, None);
let mut stream_response = self
.client
.send_chat_messages_stream(request)
.await
.map_err(|e| Self::map_ollama_error(&e))?;
let output_stream = stream! {
while let Some(chunk_result) = stream_response.next().await {
match chunk_result {
Ok(chunk) => {
let content = chunk.message.content;
if !content.is_empty() {
yield Ok(content);
}
}
Err(_) => {
yield Err(classify_stream_sse_failure("transport"));
break;
}
}
}
};
Ok(Box::new(Box::pin(output_stream)))
}
async fn stream_with_system(
&self,
system: &str,
prompt: &str,
) -> Result<Box<dyn Stream<Item = Result<String>> + Send + Unpin>> {
let messages = vec![
ChatMessage::system(system.to_string()),
ChatMessage::user(prompt.to_string()),
];
let request = self.build_request(messages, None);
let mut stream_response = self
.client
.send_chat_messages_stream(request)
.await
.map_err(|e| Self::map_ollama_error(&e))?;
let output_stream = stream! {
while let Some(chunk_result) = stream_response.next().await {
match chunk_result {
Ok(chunk) => {
let content = chunk.message.content;
if !content.is_empty() {
yield Ok(content);
}
}
Err(_) => {
yield Err(classify_stream_sse_failure("transport"));
break;
}
}
}
};
Ok(Box::new(Box::pin(output_stream)))
}
async fn stream_with_history(
&self,
messages: &[(String, String)],
) -> Result<Box<dyn Stream<Item = Result<String>> + Send + Unpin>> {
let chat_messages: Vec<ChatMessage> = messages
.iter()
.map(|(role, content)| match role.as_str() {
"system" => ChatMessage::system(content.clone()),
"user" => ChatMessage::user(content.clone()),
"assistant" => ChatMessage::assistant(content.clone()),
_ => ChatMessage::user(content.clone()),
})
.collect();
let request = self.build_request(chat_messages, None);
let mut stream_response = self
.client
.send_chat_messages_stream(request)
.await
.map_err(|e| Self::map_ollama_error(&e))?;
let output_stream = stream! {
while let Some(chunk_result) = stream_response.next().await {
match chunk_result {
Ok(chunk) => {
let content = chunk.message.content;
if !content.is_empty() {
yield Ok(content);
}
}
Err(_) => {
yield Err(classify_stream_sse_failure("transport"));
break;
}
}
}
};
Ok(Box::new(Box::pin(output_stream)))
}
fn model_name(&self) -> &str {
&self.model
}
fn supports_hints(&self) -> bool {
true
}
fn set_hints(&self, hints: GenerationHints) {
if let Ok(mut slot) = self.hints.write() {
*slot = hints;
}
}
}
impl OllamaClient {
pub async fn health_check(&self) -> Result<bool> {
match self.client.list_local_models().await {
Ok(_) => Ok(true),
Err(_) => Ok(false),
}
}
pub async fn list_models(&self) -> Result<Vec<String>> {
let models = self
.client
.list_local_models()
.await
.map_err(|e| Self::map_ollama_error(&e))?;
Ok(models.into_iter().map(|m| m.name).collect())
}
pub async fn pull_model(&self, model_name: &str) -> Result<()> {
self.client
.pull_model(model_name.to_string(), false)
.await
.map_err(|e| Self::map_ollama_error(&e))?;
Ok(())
}
pub async fn model_capabilities(&self, model_name: &str) -> Result<OllamaModelCapabilities> {
let info = self.model_info(model_name).await?;
Ok(parse_model_capabilities_from_show(&info))
}
pub async fn model_info(&self, model_name: &str) -> Result<serde_json::Value> {
let info = self
.client
.show_model_info(model_name.to_string())
.await
.map_err(|e| Self::map_ollama_error(&e))?;
Ok(serde_json::json!({
"modelfile": info.modelfile,
"parameters": info.parameters,
"template": info.template,
"capabilities": info.capabilities,
}))
}
}
#[cfg(all(test, feature = "ollama"))]
mod hint_tests {
use super::*;
use wiremock::matchers::{method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
async fn ollama_client(server: &MockServer) -> OllamaClient {
OllamaClient::with_params(
format!("http://127.0.0.1:{}", server.address().port()),
"test-model".to_string(),
ModelParams {
temperature: Some(0.4),
max_tokens: Some(64),
..ModelParams::default()
},
)
.await
.expect("client")
}
fn chat_ok() -> String {
serde_json::json!({
"model": "test-model",
"created_at": "now",
"message": { "role": "assistant", "content": "ok" },
"done": true,
"prompt_eval_count": 3,
"eval_count": 2
})
.to_string()
}
#[tokio::test]
async fn unsupported_backend_ignores_grammar_hint() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/api/chat"))
.respond_with(ResponseTemplate::new(200).set_body_string(chat_ok()))
.mount(&server)
.await;
let client = ollama_client(&server).await;
client.set_hints(GenerationHints {
json_mode: false,
suppress_reasoning: false,
max_tokens: None,
guided_grammar: Some("root ::= \"yes\" | \"no\"".to_string()),
});
let _ = client.generate("hi").await.unwrap();
let requests = server.received_requests().await.unwrap();
assert_eq!(requests.len(), 1);
let body: serde_json::Value = serde_json::from_slice(&requests[0].body).unwrap();
assert!(
body.get("guided_grammar").is_none(),
"unsupported backend must not carry the grammar field"
);
assert!(
body.get("grammar").is_none(),
"no grammar-shaped key may appear on the wire"
);
assert_eq!(body["options"]["num_predict"], 64);
}
#[tokio::test]
async fn hints_map_to_num_predict_and_json_format() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/api/chat"))
.respond_with(ResponseTemplate::new(200).set_body_string(chat_ok()))
.mount(&server)
.await;
let client = ollama_client(&server).await;
assert!(client.supports_hints());
client.set_hints(GenerationHints {
json_mode: true,
suppress_reasoning: false,
max_tokens: Some(256),
guided_grammar: None,
});
let _ = client.generate("hi").await.unwrap();
let requests = server.received_requests().await.unwrap();
let body: serde_json::Value = serde_json::from_slice(&requests[0].body).unwrap();
let options = &body["options"];
assert_eq!(
options["num_predict"], 256,
"hint max_tokens maps to num_predict"
);
assert_eq!(
options["num_ctx"],
serde_json::Value::Null,
"sanity: untouched option keys stay absent"
);
assert!(
body.get("format").is_none() || body["format"].is_null(),
"json_mode must be gated on supports_json_mode"
);
}
#[tokio::test]
async fn json_mode_applies_when_capability_supports_it() {
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/api/chat"))
.respond_with(ResponseTemplate::new(200).set_body_string(chat_ok()))
.mount(&server)
.await;
let client = ollama_client(&server).await;
client.set_hints(GenerationHints {
json_mode: true,
suppress_reasoning: false,
max_tokens: None,
guided_grammar: None,
});
let caps = parse_model_capabilities_from_show(&serde_json::json!({
"parameters": "num_ctx 2048",
"template": "",
"modelfile": "FROM x\nPARAMETER format json"
}));
assert!(
caps.supports_json_mode,
"fixture must parse as json-capable"
);
let messages = vec![ChatMessage::user("hi".into())];
let request = client.build_request(messages, Some(&caps));
let value = serde_json::to_value(&request).unwrap();
assert_eq!(value["format"], serde_json::json!("json"));
}
}