use serde_json::{json, Value};
use crate::base64;
use crate::schema::{build_json_schema, system_prompt, user_prompt};
use crate::validate::{parse_response, truncate};
use crate::{AnalyzeRequest, FieldType, ModelConfig, ModelError, Provider, Step, ToolSpec};
pub const DEFAULT_TIMEOUT_SECS: u64 = 600;
pub const DEFAULT_MAX_RETRIES: u32 = 2;
const RETRY_BASE_MS: u64 = 500;
const MAX_RETRY_AFTER_SECS: u64 = 20;
const OPENAI_BASE: &str = "https://api.openai.com/v1";
const OLLAMA_BASE: &str = "http://localhost:11434";
pub(crate) type Transport = dyn Fn(&str, &[(&str, String)], &Value) -> Result<String, ModelError>;
pub fn parse_model_spec(spec: &str) -> Result<ModelConfig, ModelError> {
let (scheme, model) = spec.split_once(':').ok_or_else(|| {
ModelError::new(format!(
"model spec `{spec}` needs a provider prefix, e.g. `openai:gpt-4o` or `local:llama3.1:8b`"
))
})?;
if model.trim().is_empty() {
return Err(ModelError::new(format!(
"model spec `{spec}` has no model name after `{scheme}:`"
)));
}
let provider = match scheme {
"openai" => Provider::OpenAI,
"local" | "ollama" => Provider::Ollama,
other => {
return Err(ModelError::new(format!(
"unknown model provider `{other}` (expected `openai` or `local`)"
)))
}
};
Ok(ModelConfig {
provider,
model: model.to_string(),
endpoint: None,
api_key: None,
max_output_tokens: 4096,
timeout_secs: DEFAULT_TIMEOUT_SECS,
max_retries: DEFAULT_MAX_RETRIES,
})
}
pub(crate) fn step_with(
config: &ModelConfig,
req: &AnalyzeRequest,
transport: &Transport,
) -> Result<Step, ModelError> {
match config.provider {
Provider::OpenAI => openai(config, req, transport),
Provider::Ollama => ollama(config, req, transport),
}
}
fn tools_json(tools: &[ToolSpec]) -> Value {
Value::Array(
tools
.iter()
.map(|tool| {
let mut properties = serde_json::Map::new();
let mut required = Vec::new();
for (name, ty) in &tool.params {
properties.insert(name.clone(), param_schema(ty));
required.push(Value::String(name.clone()));
}
json!({
"type": "function",
"function": {
"name": tool.name,
"description": tool.description,
"parameters": {
"type": "object",
"properties": properties,
"required": required,
}
}
})
})
.collect(),
)
}
fn param_schema(ty: &FieldType) -> Value {
match ty {
FieldType::Str => json!({"type": "string"}),
FieldType::Int => json!({"type": "integer"}),
FieldType::Float => json!({"type": "number"}),
FieldType::Bool => json!({"type": "boolean"}),
FieldType::ListOfStr => json!({"type": "array", "items": {"type": "string"}}),
FieldType::Object(nested) => crate::schema::object_schema(nested),
FieldType::ListOfObject(nested) => {
json!({"type": "array", "items": crate::schema::object_schema(nested)})
}
}
}
fn messages(req: &AnalyzeRequest, provider: &Provider) -> Vec<Value> {
let text = user_prompt(&req.prompt, &req.data_json);
let user = if req.images.is_empty() {
json!({"role": "user", "content": text})
} else {
match provider {
Provider::OpenAI => {
let mut parts = vec![json!({"type": "text", "text": text})];
for image in &req.images {
parts.push(json!({
"type": "image_url",
"image_url": {
"url": format!(
"data:{};base64,{}",
image.mime,
base64::encode(&image.bytes)
)
}
}));
}
json!({"role": "user", "content": parts})
}
Provider::Ollama => {
let encoded: Vec<Value> = req
.images
.iter()
.map(|i| Value::String(base64::encode(&i.bytes)))
.collect();
json!({"role": "user", "content": text, "images": encoded})
}
}
};
let mut out = vec![
json!({"role": "system", "content": system_prompt(&req.schema)}),
user,
];
for exchange in &req.tool_history {
out.push(json!({
"role": "assistant",
"content": format!("Calling {}({})", exchange.name, exchange.arguments_json),
}));
out.push(json!({
"role": "user",
"content": format!("Result of {}: {}", exchange.name, exchange.result_json),
}));
}
out
}
fn openai(
config: &ModelConfig,
req: &AnalyzeRequest,
transport: &Transport,
) -> Result<Step, ModelError> {
let key = config
.api_key
.clone()
.or_else(|| std::env::var("OPENAI_API_KEY").ok())
.filter(|k| !k.trim().is_empty())
.ok_or_else(|| {
ModelError::new("OPENAI_API_KEY not set (export it, or set api_key in kora.toml)")
})?;
let mut body = json!({
"model": config.model,
"max_completion_tokens": config.max_output_tokens,
"messages": messages(req, &Provider::OpenAI),
});
if req.tools.is_empty() {
body["response_format"] = json!({
"type": "json_schema",
"json_schema": {
"name": sanitize_schema_name(&req.schema.type_name),
"strict": true,
"schema": build_json_schema(&req.schema),
}
});
} else {
body["tools"] = tools_json(&req.tools);
}
let headers = [
("Authorization", format!("Bearer {key}")),
("Content-Type", "application/json".to_string()),
];
let url = format!("{OPENAI_BASE}/chat/completions");
let text = transport(&url, &headers, &body)?;
let response: Value = serde_json::from_str(&text).map_err(|e| {
ModelError::new(format!(
"OpenAI returned a non-JSON body ({e}): {}",
truncate(&text, 300)
))
})?;
let tokens_in = response["usage"]["prompt_tokens"].as_u64().unwrap_or(0);
let tokens_out = response["usage"]["completion_tokens"].as_u64().unwrap_or(0);
let message = &response["choices"][0]["message"];
if let Some(call) = message["tool_calls"].get(0) {
let name = call["function"]["name"].as_str().unwrap_or_default();
let arguments_json = call["function"]["arguments"]
.as_str()
.unwrap_or("{}")
.to_string();
return Ok(Step::CallTool {
name: name.to_string(),
arguments_json,
tokens_in,
tokens_out,
});
}
let content = message["content"].as_str().ok_or_else(|| {
ModelError::new(format!(
"OpenAI response had no message content: {}",
truncate(&text, 300)
))
})?;
parse_response(content, &req.schema, tokens_in, tokens_out).map(Step::Done)
}
fn ollama(
config: &ModelConfig,
req: &AnalyzeRequest,
transport: &Transport,
) -> Result<Step, ModelError> {
let base = config.endpoint.as_deref().unwrap_or(OLLAMA_BASE);
let mut body = json!({
"model": config.model,
"stream": false,
"options": {"num_predict": config.max_output_tokens},
"messages": messages(req, &Provider::Ollama),
});
if req.tools.is_empty() {
body["format"] = build_json_schema(&req.schema);
} else {
body["tools"] = tools_json(&req.tools);
}
let headers = [("Content-Type", "application/json".to_string())];
let url = format!("{}/api/chat", base.trim_end_matches('/'));
let text = transport(&url, &headers, &body)?;
let response: Value = serde_json::from_str(&text).map_err(|e| {
ModelError::new(format!(
"Ollama returned a non-JSON body ({e}): {}",
truncate(&text, 300)
))
})?;
let tokens_in = response["prompt_eval_count"].as_u64().unwrap_or(0);
let tokens_out = response["eval_count"].as_u64().unwrap_or(0);
let message = &response["message"];
if let Some(call) = message["tool_calls"].get(0) {
let name = call["function"]["name"].as_str().unwrap_or_default();
let arguments_json = match &call["function"]["arguments"] {
Value::String(s) => s.clone(),
other => other.to_string(),
};
return Ok(Step::CallTool {
name: name.to_string(),
arguments_json,
tokens_in,
tokens_out,
});
}
let content = message["content"].as_str().ok_or_else(|| {
ModelError::new(format!(
"Ollama response had no message content: {}",
truncate(&text, 300)
))
})?;
parse_response(content, &req.schema, tokens_in, tokens_out).map(Step::Done)
}
fn sanitize_schema_name(name: &str) -> String {
let cleaned: String = name
.chars()
.map(|c| {
if c.is_ascii_alphanumeric() || c == '_' || c == '-' {
c
} else {
'_'
}
})
.collect();
if cleaned.is_empty() {
"Result".to_string()
} else {
cleaned
}
}
pub(crate) fn transport_for(config: &ModelConfig) -> Box<Transport> {
let timeout = std::time::Duration::from_secs(config.timeout_secs.max(1));
let attempts = config.max_retries.saturating_add(1);
Box::new(move |url: &str, headers: &[(&str, String)], body: &Value| {
retry_loop(attempts, || send(url, headers, body, timeout))
})
}
fn retry_loop<F>(attempts: u32, mut attempt_once: F) -> Result<String, ModelError>
where
F: FnMut() -> Result<String, (ModelError, Option<u64>)>,
{
let mut attempt = 0;
loop {
attempt += 1;
let (error, retry_after) = match attempt_once() {
Ok(text) => return Ok(text),
Err(e) => e,
};
if !error.retryable || attempt >= attempts {
return Err(error);
}
std::thread::sleep(retry_delay(attempt, retry_after));
}
}
fn retry_delay(attempt: u32, retry_after: Option<u64>) -> std::time::Duration {
if let Some(secs) = retry_after {
return std::time::Duration::from_secs(secs.min(MAX_RETRY_AFTER_SECS));
}
let base = RETRY_BASE_MS.saturating_mul(1 << (attempt - 1).min(5));
std::time::Duration::from_millis(base + jitter_ms(base))
}
fn jitter_ms(base: u64) -> u64 {
let nanos = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map(|d| d.subsec_nanos() as u64)
.unwrap_or(0);
nanos % (base / 4).max(1)
}
fn retryable_status(code: u16) -> bool {
matches!(code, 408 | 409 | 429) || (500..600).contains(&code)
}
fn retry_after_secs(response: &ureq::Response) -> Option<u64> {
response.header("retry-after")?.trim().parse::<u64>().ok()
}
#[allow(clippy::type_complexity)]
fn send(
url: &str,
headers: &[(&str, String)],
body: &Value,
timeout: std::time::Duration,
) -> Result<String, (ModelError, Option<u64>)> {
let agent = ureq::AgentBuilder::new().timeout(timeout).build();
let mut request = agent.post(url);
for (name, value) in headers {
request = request.set(name, value);
}
match request.send_json(body.clone()) {
Ok(response) => response.into_string().map_err(|e| {
(
ModelError::retryable(format!("could not read response body from {url}: {e}")),
None,
)
}),
Err(ureq::Error::Status(code, response)) => {
let retry_after = retry_after_secs(&response);
let body = response.into_string().unwrap_or_default();
let message = format!("{url} returned HTTP {code}: {}", truncate(&body, 300));
let error = if retryable_status(code) {
ModelError::retryable(message)
} else {
ModelError::new(message)
};
Err((error, retry_after))
}
Err(e) => Err((
ModelError::retryable(format!("request to {url} failed: {e}")),
None,
)),
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::{AnalyzeOutcome, FieldType, Schema, SchemaField};
use std::cell::RefCell;
fn schema() -> Schema {
Schema {
type_name: "Insight".into(),
fields: vec![
SchemaField {
name: "summary".into(),
field_type: FieldType::Str,
description: None,
pattern: None,
},
SchemaField {
name: "count".into(),
field_type: FieldType::Int,
description: None,
pattern: None,
},
],
}
}
fn request() -> AnalyzeRequest {
AnalyzeRequest {
prompt: "find anomalies".into(),
data_json: "{\"rows\":2}".into(),
images: Vec::new(),
schema: schema(),
tools: Vec::new(),
tool_history: Vec::new(),
}
}
type Captured = std::rc::Rc<RefCell<Option<(String, Value)>>>;
type Recorder = (Box<Transport>, Captured);
fn recording(reply: &'static str) -> Recorder {
let seen: Captured = std::rc::Rc::new(RefCell::new(None));
let sink = seen.clone();
let transport = Box::new(move |url: &str, _h: &[(&str, String)], body: &Value| {
*sink.borrow_mut() = Some((url.to_string(), body.clone()));
Ok(reply.to_string())
});
(transport, seen)
}
#[test]
fn spec_openai() {
let c = parse_model_spec("openai:gpt-4o").unwrap();
assert_eq!(c.provider, Provider::OpenAI);
assert_eq!(c.model, "gpt-4o");
assert_eq!(c.max_output_tokens, 4096);
assert_eq!(c.timeout_secs, DEFAULT_TIMEOUT_SECS);
}
#[test]
fn spec_local_keeps_tag() {
let c = parse_model_spec("local:llama3.1:8b").unwrap();
assert_eq!(c.provider, Provider::Ollama);
assert_eq!(c.model, "llama3.1:8b");
}
#[test]
fn spec_errors() {
assert!(parse_model_spec("gpt-4o")
.unwrap_err()
.message
.contains("prefix"));
assert!(parse_model_spec("openai:")
.unwrap_err()
.message
.contains("no model name"));
assert!(parse_model_spec("groq:x")
.unwrap_err()
.message
.contains("unknown model provider"));
}
#[test]
fn openai_request_shape_and_parse() {
let reply = r#"{
"choices":[{"message":{"content":"{\"summary\":\"ok\",\"count\":2,\"__uncertain__\":\"\"}"}}],
"usage":{"prompt_tokens":11,"completion_tokens":7}
}"#;
let (transport, seen) = recording(reply);
let mut config = parse_model_spec("openai:gpt-4o").unwrap();
config.api_key = Some("test-key".into());
let outcome = step_with(&config, &request(), &*transport).unwrap();
match outcome {
Step::Done(AnalyzeOutcome::Ok {
fields_json,
tokens_in,
tokens_out,
}) => {
assert_eq!(fields_json["summary"], "ok");
assert_eq!(tokens_in, 11);
assert_eq!(tokens_out, 7);
}
other => panic!("expected Ok, got {other:?}"),
}
let (url, body) = seen.borrow().clone().unwrap();
assert_eq!(url, "https://api.openai.com/v1/chat/completions");
assert_eq!(body["response_format"]["type"], "json_schema");
assert_eq!(body["response_format"]["json_schema"]["strict"], true);
assert_eq!(body["messages"][0]["role"], "system");
assert!(body["messages"][1]["content"]
.as_str()
.unwrap()
.contains("DATA:"));
}
#[test]
fn openai_missing_key_is_clear() {
std::env::remove_var("OPENAI_API_KEY");
let (transport, _seen) = recording("{}");
let config = parse_model_spec("openai:gpt-4o").unwrap();
let err = step_with(&config, &request(), &*transport).unwrap_err();
assert!(
err.message.contains("OPENAI_API_KEY not set"),
"{}",
err.message
);
}
#[test]
fn ollama_request_shape_and_uncertain() {
let reply = r#"{
"message":{"content":"{\"summary\":\"\",\"count\":0,\"__uncertain__\":\"no revenue column\"}"},
"prompt_eval_count":30,"eval_count":9
}"#;
let (transport, seen) = recording(reply);
let config = parse_model_spec("local:llama3.1:8b").unwrap();
match step_with(&config, &request(), &*transport).unwrap() {
Step::Done(AnalyzeOutcome::Uncertain {
reason,
tokens_in,
tokens_out,
}) => {
assert_eq!(reason, "no revenue column");
assert_eq!(tokens_in, 30);
assert_eq!(tokens_out, 9);
}
other => panic!("expected Uncertain, got {other:?}"),
}
let (url, body) = seen.borrow().clone().unwrap();
assert_eq!(url, "http://localhost:11434/api/chat");
assert_eq!(body["stream"], false);
assert_eq!(body["format"]["type"], "object");
}
#[test]
fn ollama_endpoint_override() {
let reply =
r#"{"message":{"content":"{\"summary\":\"a\",\"count\":1,\"__uncertain__\":\"\"}"}}"#;
let (transport, seen) = recording(reply);
let mut config = parse_model_spec("local:llama3.1:8b").unwrap();
config.endpoint = Some("http://box:11434/".into());
step_with(&config, &request(), &*transport).unwrap();
assert_eq!(
seen.borrow().clone().unwrap().0,
"http://box:11434/api/chat"
);
}
#[test]
fn openai_attaches_images_as_content_parts() {
let reply = r#"{
"choices":[{"message":{"content":"{\"summary\":\"ok\",\"count\":1,\"__uncertain__\":\"\"}"}}],
"usage":{"prompt_tokens":1,"completion_tokens":1}
}"#;
let (transport, seen) = recording(reply);
let mut config = parse_model_spec("openai:gpt-4o").unwrap();
config.api_key = Some("test-key".into());
let mut req = request();
req.images = vec![crate::ImagePart {
mime: "image/png".into(),
bytes: b"foobar".to_vec(),
}];
step_with(&config, &req, &*transport).unwrap();
let (_, body) = seen.borrow().clone().unwrap();
let parts = &body["messages"][1]["content"];
assert_eq!(parts[0]["type"], "text");
assert_eq!(parts[1]["type"], "image_url");
assert_eq!(
parts[1]["image_url"]["url"],
"data:image/png;base64,Zm9vYmFy"
);
}
#[test]
fn ollama_attaches_images_beside_the_text() {
let reply =
r#"{"message":{"content":"{\"summary\":\"a\",\"count\":1,\"__uncertain__\":\"\"}"}}"#;
let (transport, seen) = recording(reply);
let config = parse_model_spec("local:llava:7b").unwrap();
let mut req = request();
req.images = vec![crate::ImagePart {
mime: "image/png".into(),
bytes: b"foobar".to_vec(),
}];
step_with(&config, &req, &*transport).unwrap();
let (_, body) = seen.borrow().clone().unwrap();
let message = &body["messages"][1];
assert!(message["content"].as_str().unwrap().contains("DATA:"));
assert_eq!(message["images"][0], "Zm9vYmFy");
}
#[test]
fn no_images_keeps_plain_string_content() {
let reply = r#"{
"choices":[{"message":{"content":"{\"summary\":\"ok\",\"count\":1,\"__uncertain__\":\"\"}"}}]
}"#;
let (transport, seen) = recording(reply);
let mut config = parse_model_spec("openai:gpt-4o").unwrap();
config.api_key = Some("test-key".into());
step_with(&config, &request(), &*transport).unwrap();
let (_, body) = seen.borrow().clone().unwrap();
assert!(body["messages"][1]["content"].is_string());
}
#[test]
fn schema_name_sanitized() {
assert_eq!(sanitize_schema_name("Insight"), "Insight");
assert_eq!(sanitize_schema_name("my type!"), "my_type_");
assert_eq!(sanitize_schema_name(""), "Result");
}
}
#[cfg(test)]
mod retry_tests {
use super::*;
#[test]
fn a_bad_request_is_never_retried() {
for code in [400, 401, 403, 404, 422] {
assert!(!retryable_status(code), "{code} should not be retried");
}
}
#[test]
fn a_rate_limit_or_a_server_error_is_retried() {
for code in [408, 409, 429, 500, 502, 503, 504] {
assert!(retryable_status(code), "{code} should be retried");
}
}
#[test]
fn the_backoff_grows_and_stays_bounded() {
let first = retry_delay(1, None);
let second = retry_delay(2, None);
assert!(
first.as_millis() >= RETRY_BASE_MS as u128,
"the first wait should be at least the base"
);
assert!(
second >= first,
"waits should grow: {second:?} came after {first:?}"
);
assert!(first.as_millis() < (RETRY_BASE_MS as u128) * 2);
}
#[test]
fn the_provider_is_believed_over_the_local_backoff() {
assert_eq!(retry_delay(1, Some(3)).as_secs(), 3);
}
#[test]
fn an_absurd_retry_after_is_capped_rather_than_waited_out() {
assert_eq!(
retry_delay(1, Some(3600)).as_secs(),
MAX_RETRY_AFTER_SECS,
"a long Retry-After should be capped"
);
}
#[test]
fn retries_stop_at_the_configured_count() {
let attempts = std::cell::Cell::new(0);
let result = retry_loop(3, || {
attempts.set(attempts.get() + 1);
Err((ModelError::retryable("nope"), Some(0)))
});
assert!(result.is_err());
assert_eq!(attempts.get(), 3, "three attempts, then give up");
}
#[test]
fn an_unretryable_failure_is_reported_on_the_first_attempt() {
let attempts = std::cell::Cell::new(0);
let result = retry_loop(3, || {
attempts.set(attempts.get() + 1);
Err((ModelError::new("bad api key"), None))
});
assert!(result.is_err());
assert_eq!(attempts.get(), 1, "a 401 does not improve with waiting");
}
#[test]
fn a_retry_that_succeeds_returns_the_answer() {
let attempts = std::cell::Cell::new(0);
let result = retry_loop(3, || {
attempts.set(attempts.get() + 1);
if attempts.get() < 2 {
Err((ModelError::retryable("try again"), Some(0)))
} else {
Ok("body".to_string())
}
});
assert_eq!(result.unwrap(), "body");
assert_eq!(attempts.get(), 2);
}
}