kibble 0.1.0

chew through any source into clean datasets — a fast ingestion, RAG & fine-tuning toolkit
Documentation
//! Shared OpenAI-compatible chat helper (used by `bench` and `ask`).

use std::time::Instant;

/// A tool/function call the model requested.
#[derive(Debug)]
pub struct ToolCall {
    pub id: String,
    pub name: String,
    pub arguments: String,
}

/// Sampling parameters for one chat turn. `top_p`/`top_k`/`repetition_penalty`
/// are sent on every request: plain temperature sampling drives thinking models
/// (Qwen3.5 / Qwythos) into repetition loops, so nucleus+top-k plus a light
/// repetition penalty is the model's designed operating point. Servers that
/// don't recognise `top_k`/`repetition_penalty` (strict OpenAI) ignore them.
#[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 {
    /// Recommended sampling for Qwen3.5-family models at the given
    /// temperature/token budget (top_p 0.95, top_k 20, repetition_penalty 1.05).
    pub fn recommended(temperature: f64, max_tokens: usize) -> Self {
        Sampling { temperature, max_tokens, top_p: 0.95, top_k: 20, repetition_penalty: 1.05 }
    }
}

/// Extract an error message when the response body is an error object rather than a
/// completion. Servers sometimes return HTTP 200 with `{"error": ...}` (e.g. mlx_lm.server
/// wrapping a HuggingFace model-not-found), which must surface as an error — otherwise it
/// looks like an empty completion and gets misread downstream as "no answer" / NOTFOUND.
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(),
    )
}

/// One `/chat/completions` turn. Returns (message content, tool calls,
/// completion tokens, elapsed ms). Sends `tools` with `tool_choice:"auto"` when present.
#[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))
}

/// Streaming `/chat/completions` turn: sends `stream:true`, invokes `on_content`
/// with each content delta as it arrives, and returns the accumulated
/// `(content, tool_calls)` — the same shape as `chat_turn`.
#[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();
    // index -> (id, name, accumulated arguments)
    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);
                    }
                }
            }
        }
    }

    // A non-SSE error body (e.g. mlx wrapping an HF 404 at HTTP 200) never matches a
    // `data:` line, so it sits unparsed in `buf`. If nothing streamed, surface it.
    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))
}

/// First non-empty env var among `names`, else empty string.
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() {
        // The request body must carry top_p/top_k/repetition_penalty — without them,
        // thinking models loop. Capture the raw request and assert the fields are present.
        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));
    }

    // Serve one canned HTTP response body on a fresh loopback port; returns "http://127.0.0.1:PORT".
    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() {
        // Content in two frames, then a tool_call split across two frames, then DONE.
        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\"}"); // reassembled across frames
    }

    #[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() {
        // Some servers (mlx_lm.server wrapping a bad model id) return HTTP 200 with an
        // {"error": ...} body. That must become an Err, not a silent empty completion
        // (which downstream misreads as "no answer" / NOTFOUND).
        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() {
        // Same, over the streaming path: a non-SSE error body must Err, not return empty.
        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}");
    }
}