use std::time::Instant;
#[derive(Debug)]
pub struct ToolCall {
pub id: String,
pub name: String,
pub arguments: String,
}
#[derive(Clone, Copy)]
pub struct Sampling {
pub temperature: f64,
pub max_tokens: usize,
pub top_p: f64,
pub top_k: usize,
pub repetition_penalty: f64,
}
impl Sampling {
pub fn recommended(temperature: f64, max_tokens: usize) -> Self {
Sampling { temperature, max_tokens, top_p: 0.95, top_k: 20, repetition_penalty: 1.05 }
}
}
fn server_error(v: &serde_json::Value) -> Option<String> {
let e = v.get("error")?;
Some(
e.get("message").and_then(|m| m.as_str())
.or_else(|| e.as_str())
.unwrap_or("unknown error")
.to_string(),
)
}
#[allow(clippy::too_many_arguments)]
pub async fn chat_turn(
client: &reqwest::Client, base_url: &str, model: &str, key: &str,
messages: &[serde_json::Value], tools: Option<&serde_json::Value>,
sampling: Sampling,
) -> std::io::Result<(Option<String>, Vec<ToolCall>, usize, u128)> {
let mut body = serde_json::json!({"model":model,"messages":messages,
"temperature":sampling.temperature,"max_tokens":sampling.max_tokens,
"top_p":sampling.top_p,"top_k":sampling.top_k,"repetition_penalty":sampling.repetition_penalty});
if let Some(t) = tools {
body["tools"] = t.clone();
body["tool_choice"] = serde_json::json!("auto");
}
let url = format!("{}/chat/completions", base_url.trim_end_matches('/'));
let start = Instant::now();
let mut req = client.post(&url).header("content-type","application/json").body(body.to_string());
if !key.is_empty() { req = req.header("authorization", format!("Bearer {key}")); }
let resp = req.send().await.map_err(|e| std::io::Error::other(format!("chat request: {e}")))?;
let text = resp.text().await.map_err(|e| std::io::Error::other(format!("chat body: {e}")))?;
let elapsed = start.elapsed().as_millis();
let v: serde_json::Value = serde_json::from_str(&text)
.map_err(|e| std::io::Error::new(std::io::ErrorKind::InvalidData, format!("chat json: {e}")))?;
if let Some(msg) = server_error(&v) {
return Err(std::io::Error::other(format!("LLM server error: {msg}")));
}
let msg = &v["choices"][0]["message"];
let content = msg.get("content").and_then(|c| c.as_str()).filter(|s| !s.is_empty()).map(String::from);
let mut calls = Vec::new();
if let Some(tcs) = msg.get("tool_calls").and_then(|t| t.as_array()) {
for tc in tcs {
calls.push(ToolCall {
id: tc.get("id").and_then(|x| x.as_str()).unwrap_or("").to_string(),
name: tc["function"]["name"].as_str().unwrap_or("").to_string(),
arguments: tc["function"]["arguments"].as_str().unwrap_or("{}").to_string(),
});
}
}
let tokens = v["usage"]["completion_tokens"].as_u64()
.unwrap_or_else(|| content.as_deref().unwrap_or("").split_whitespace().count() as u64) as usize;
Ok((content, calls, tokens, elapsed))
}
#[allow(clippy::too_many_arguments)]
pub async fn chat_turn_stream<F: FnMut(&str)>(
client: &reqwest::Client, base_url: &str, model: &str, key: &str,
messages: &[serde_json::Value], tools: Option<&serde_json::Value>,
sampling: Sampling, mut on_content: F,
) -> std::io::Result<(Option<String>, Vec<ToolCall>)> {
let mut body = serde_json::json!({"model":model,"messages":messages,
"temperature":sampling.temperature,"max_tokens":sampling.max_tokens,
"top_p":sampling.top_p,"top_k":sampling.top_k,"repetition_penalty":sampling.repetition_penalty,"stream":true});
if let Some(t) = tools {
body["tools"] = t.clone();
body["tool_choice"] = serde_json::json!("auto");
}
let url = format!("{}/chat/completions", base_url.trim_end_matches('/'));
let mut req = client.post(&url).header("content-type", "application/json").body(body.to_string());
if !key.is_empty() {
req = req.header("authorization", format!("Bearer {key}"));
}
let mut resp = req.send().await.and_then(|r| r.error_for_status())
.map_err(|e| std::io::Error::other(format!("chat stream request: {e}")))?;
let mut buf: Vec<u8> = Vec::new();
let mut content = String::new();
let mut tcs: std::collections::BTreeMap<usize, (String, String, String)> = std::collections::BTreeMap::new();
'outer: while let Some(chunk) = resp.chunk().await.map_err(|e| std::io::Error::other(format!("chat stream body: {e}")))? {
buf.extend_from_slice(&chunk);
while let Some(pos) = buf.iter().position(|&b| b == b'\n') {
let line_bytes: Vec<u8> = buf.drain(..=pos).collect();
let line_cow = String::from_utf8_lossy(&line_bytes);
let line = line_cow.trim();
let Some(data) = line.strip_prefix("data:") else { continue };
let data = data.trim();
if data == "[DONE]" {
break 'outer;
}
if data.is_empty() {
continue;
}
let Ok(v) = serde_json::from_str::<serde_json::Value>(data) else { continue };
let delta = &v["choices"][0]["delta"];
if let Some(piece) = delta.get("content").and_then(|c| c.as_str()) {
if !piece.is_empty() {
content.push_str(piece);
on_content(piece);
}
}
if let Some(arr) = delta.get("tool_calls").and_then(|t| t.as_array()) {
for c in arr {
let idx = c.get("index").and_then(|i| i.as_u64()).unwrap_or(0) as usize;
let e = tcs.entry(idx).or_default();
if let Some(id) = c.get("id").and_then(|x| x.as_str()) {
if !id.is_empty() { e.0 = id.to_string(); }
}
let f = &c["function"];
if let Some(name) = f.get("name").and_then(|x| x.as_str()) {
if !name.is_empty() { e.1 = name.to_string(); }
}
if let Some(args) = f.get("arguments").and_then(|x| x.as_str()) {
e.2.push_str(args);
}
}
}
}
}
if content.is_empty() && tcs.is_empty() {
let leftover = String::from_utf8_lossy(&buf);
let leftover = leftover.trim();
if !leftover.is_empty() {
if let Ok(v) = serde_json::from_str::<serde_json::Value>(leftover) {
if let Some(msg) = server_error(&v) {
return Err(std::io::Error::other(format!("LLM server error: {msg}")));
}
}
}
}
let tool_calls: Vec<ToolCall> = tcs.into_values()
.map(|(id, name, arguments)| ToolCall { id, name, arguments })
.collect();
Ok((if content.is_empty() { None } else { Some(content) }, tool_calls))
}
pub fn env_key(names: &[&str]) -> String {
for n in names {
if let Ok(v) = std::env::var(n) {
if !v.is_empty() { return v; }
}
}
String::new()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn env_key_picks_first_present() {
std::env::set_var("KIBBLE_LLM_TEST_A", "");
std::env::set_var("KIBBLE_LLM_TEST_B", "beta");
assert_eq!(env_key(&["KIBBLE_LLM_TEST_A", "KIBBLE_LLM_TEST_B"]), "beta");
assert_eq!(env_key(&["KIBBLE_LLM_TEST_MISSING"]), "");
std::env::remove_var("KIBBLE_LLM_TEST_A");
std::env::remove_var("KIBBLE_LLM_TEST_B");
}
#[tokio::test]
async fn chat_turn_parses_content_and_tool_calls() {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let server = tokio::spawn(async move {
let (mut s, _) = listener.accept().await.unwrap();
use tokio::io::{AsyncReadExt, AsyncWriteExt};
let mut b = [0u8; 4096]; let _ = s.read(&mut b).await.unwrap();
let json = r#"{"choices":[{"message":{"content":null,"tool_calls":[{"id":"c1","type":"function","function":{"name":"search_corpus","arguments":"{\"query\":\"fox\"}"}}]}}],"usage":{"completion_tokens":3}}"#;
let resp = format!("HTTP/1.1 200 OK\r\nContent-Type: application/json\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{}", json.len(), json);
s.write_all(resp.as_bytes()).await.unwrap();
});
let client = reqwest::Client::new();
let (content, calls, _tok, _ms) = chat_turn(&client, &format!("http://{addr}/v1"), "m", "", &[serde_json::json!({"role":"user","content":"hi"})], None, Sampling::recommended(0.0, 64)).await.unwrap();
server.await.unwrap();
assert!(content.is_none());
assert_eq!(calls.len(), 1);
assert_eq!(calls[0].name, "search_corpus");
assert!(calls[0].arguments.contains("fox"));
}
#[tokio::test]
async fn chat_turn_sends_sampling_params() {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let server = tokio::spawn(async move {
let (mut s, _) = listener.accept().await.unwrap();
use tokio::io::{AsyncReadExt, AsyncWriteExt};
let mut b = vec![0u8; 8192];
let n = s.read(&mut b).await.unwrap();
let req = String::from_utf8_lossy(&b[..n]).to_string();
let json = r#"{"choices":[{"message":{"content":"ok"}}],"usage":{"completion_tokens":1}}"#;
let resp = format!("HTTP/1.1 200 OK\r\nContent-Type: application/json\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{}", json.len(), json);
s.write_all(resp.as_bytes()).await.unwrap();
req
});
let client = reqwest::Client::new();
let sampling = Sampling { temperature: 0.6, max_tokens: 128, top_p: 0.9, top_k: 25, repetition_penalty: 1.07 };
let _ = chat_turn(&client, &format!("http://{addr}/v1"), "m", "", &[serde_json::json!({"role":"user","content":"hi"})], None, sampling).await.unwrap();
let req = server.await.unwrap();
let body = req.split("\r\n\r\n").nth(1).expect("request body");
let v: serde_json::Value = serde_json::from_str(body).expect("body is json");
assert_eq!(v["top_p"], serde_json::json!(0.9));
assert_eq!(v["top_k"], serde_json::json!(25));
assert_eq!(v["repetition_penalty"], serde_json::json!(1.07));
}
fn mock_sse(body: String) -> String {
let listener = std::net::TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
std::thread::spawn(move || {
use std::io::{Read, Write};
if let Ok((mut s, _)) = listener.accept() {
let mut b = [0u8; 4096];
let _ = s.read(&mut b);
let resp = format!("HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{}", body.len(), body);
let _ = s.write_all(resp.as_bytes());
}
});
format!("http://{addr}")
}
#[tokio::test]
async fn chat_turn_stream_accumulates_content_and_toolcalls() {
let body = concat!(
"data: {\"choices\":[{\"delta\":{\"content\":\"Hello\"}}]}\n\n",
"data: {\"choices\":[{\"delta\":{\"content\":\" world\"}}]}\n\n",
"data: {\"choices\":[{\"delta\":{\"tool_calls\":[{\"index\":0,\"id\":\"c1\",\"function\":{\"name\":\"search_corpus\",\"arguments\":\"{\\\"que\"}}]}}]}\n\n",
"data: {\"choices\":[{\"delta\":{\"tool_calls\":[{\"index\":0,\"function\":{\"arguments\":\"ry\\\":\\\"x\\\"}\"}}]}}]}\n\n",
"data: [DONE]\n\n",
).to_string();
let base = mock_sse(body);
let client = reqwest::Client::new();
let mut pieces: Vec<String> = Vec::new();
let (content, calls) = chat_turn_stream(
&client, &base, "m", "", &[serde_json::json!({"role":"user","content":"hi"})], None, Sampling::recommended(0.0, 64),
|p| pieces.push(p.to_string()),
).await.unwrap();
assert_eq!(pieces, vec!["Hello".to_string(), " world".to_string()]);
assert_eq!(content.as_deref(), Some("Hello world"));
assert_eq!(calls.len(), 1);
assert_eq!(calls[0].name, "search_corpus");
assert_eq!(calls[0].id, "c1");
assert_eq!(calls[0].arguments, "{\"query\":\"x\"}"); }
#[tokio::test]
async fn chat_turn_stream_preserves_multibyte() {
let body = concat!(
"data: {\"choices\":[{\"delta\":{\"content\":\"café \"}}]}\n\n",
"data: {\"choices\":[{\"delta\":{\"content\":\"日本語\"}}]}\n\n",
"data: [DONE]\n\n",
).to_string();
let base = mock_sse(body);
let client = reqwest::Client::new();
let (content, _) = chat_turn_stream(&client, &base, "m", "", &[serde_json::json!({"role":"user","content":"x"})], None, Sampling::recommended(0.0, 64), |_| {}).await.unwrap();
assert_eq!(content.as_deref(), Some("café 日本語"));
}
#[tokio::test]
async fn chat_turn_surfaces_server_error() {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let server = tokio::spawn(async move {
let (mut s, _) = listener.accept().await.unwrap();
use tokio::io::{AsyncReadExt, AsyncWriteExt};
let mut b = [0u8; 4096]; let _ = s.read(&mut b).await.unwrap();
let json = r#"{"error":"404 Client Error: Repository Not Found for kuro"}"#;
let resp = format!("HTTP/1.1 200 OK\r\nContent-Type: application/json\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{}", json.len(), json);
s.write_all(resp.as_bytes()).await.unwrap();
});
let client = reqwest::Client::new();
let r = chat_turn(&client, &format!("http://{addr}/v1"), "kuro", "", &[serde_json::json!({"role":"user","content":"hi"})], None, Sampling::recommended(0.0, 64)).await;
server.await.unwrap();
let err = r.expect_err("server error body must surface as Err, not an empty completion");
assert!(err.to_string().contains("404"), "unexpected error: {err}");
}
#[tokio::test]
async fn chat_turn_stream_surfaces_server_error() {
let base = mock_sse(r#"{"error":"404 Client Error: Repository Not Found for kuro"}"#.to_string());
let client = reqwest::Client::new();
let r = chat_turn_stream(&client, &base, "kuro", "", &[serde_json::json!({"role":"user","content":"x"})], None, Sampling::recommended(0.0, 64), |_| {}).await;
let err = r.expect_err("streamed server error must surface as Err");
assert!(err.to_string().contains("404"), "unexpected error: {err}");
}
}