diff --git a/src-tauri/src/commands/llm_commands.rs b/src-tauri/src/commands/llm_commands.rs index 3c8d7a9..18af168 100644 --- a/src-tauri/src/commands/llm_commands.rs +++ b/src-tauri/src/commands/llm_commands.rs @@ -63,17 +63,16 @@ pub async fn generate_stream( emit_busy(&app, true); tauri::async_runtime::spawn(async move { - // For now, we do a non-streaming call and emit the full response as one token - // Real SSE streaming from Ollama/OpenAI can be added later + // Real token streaming: Ollama NDJSON / OpenAI SSE, one LlmEvent::Token + // per piece as it arrives. Falls back to Done carrying the full text. let result = if llm::is_ollama(&config.provider, &config.api_url) { - call_ollama(&client, &config, &messages, temperature, max_tokens).await + call_ollama_stream(&client, &config, &messages, temperature, max_tokens, &channel).await } else { - call_openai(&client, &config, &messages, temperature, max_tokens).await + call_openai_stream(&client, &config, &messages, temperature, max_tokens, &channel).await }; match result { Ok(text) => { - let _ = channel.send(LlmEvent::Token(text.clone())); let _ = channel.send(LlmEvent::Done(text)); } Err(e) => { @@ -304,4 +303,177 @@ async fn call_openai( .first() .map(|c| c.message.content.clone()) .ok_or_else(|| "No response from API".to_string()) +} + +// ─── Real streaming (Ollama NDJSON / OpenAI SSE) ───────────── +// ponytail: shared line-buffer + per-line Value parsing. Typed structs would +// break on the shape variance across servers (final lines omit message, +// empty contents, metrics lines) — skipping unparseable lines is the +// boring-correct choice. + +use futures_util::StreamExt; + +/// What one parsed stream line yields. +struct StreamLine { + token: Option, + done: bool, +} + +/// Parse one Ollama NDJSON line: `{"message":{"content":"…"},"done":false}`. +fn parse_ollama_line(line: &str) -> Result, String> { + let line = line.trim(); + if line.is_empty() { + return Ok(None); + } + let v: serde_json::Value = serde_json::from_str(line).map_err(|e| format!("bad NDJSON line: {e}"))?; + if let Some(err) = v["error"].as_str() { + return Err(err.to_string()); + } + Ok(Some(StreamLine { + token: v["message"]["content"].as_str().filter(|s| !s.is_empty()).map(|s| s.to_string()), + done: v["done"].as_bool().unwrap_or(false), + })) +} + +/// Parse one OpenAI SSE line: `data: {"choices":[{"delta":{"content":"…"}}]}` +/// or the terminal `data: [DONE]`. +fn parse_openai_line(line: &str) -> Result, String> { + let line = line.trim(); + let Some(payload) = line.strip_prefix("data: ") else { + return Ok(None); // keep-alive comments / blank lines + }; + if payload == "[DONE]" { + return Ok(Some(StreamLine { token: None, done: true })); + } + let v: serde_json::Value = serde_json::from_str(payload).map_err(|e| format!("bad SSE line: {e}"))?; + if let Some(err) = v["error"]["message"].as_str() { + return Err(err.to_string()); + } + Ok(Some(StreamLine { + token: v["choices"][0]["delta"]["content"].as_str().filter(|s| !s.is_empty()).map(|s| s.to_string()), + done: false, + })) +} + +/// POST with `stream: true`, emit Token per piece, return the full text +/// when the stream finishes. `body` is the request JSON minus the stream flag. +async fn stream_chat( + client: &reqwest::Client, + url: &str, + mut body: serde_json::Value, + channel: &Channel, + parse: fn(&str) -> Result, String>, + bearer: Option<&str>, +) -> Result { + body["stream"] = serde_json::json!(true); + let mut req = client.post(url).json(&body); + if let Some(key) = bearer { + req = req.bearer_auth(key); + } + let res = req.send().await.map_err(|e| format!("stream request failed: {e}"))?; + if !res.status().is_success() { + let status = res.status(); + let text = res.text().await.unwrap_or_default(); + return Err(format!("stream error {status}: {text}")); + } + stream_from_response(res, channel, parse).await +} + +/// Consume a streaming HTTP response: buffer bytes into lines, parse each +/// with `parse`, emit Token per piece, Done via the caller's Ok(full text). +async fn stream_from_response( + res: reqwest::Response, + channel: &Channel, + parse: fn(&str) -> Result, String>, +) -> Result { + let mut full = String::new(); + let mut buf = String::new(); + let mut stream = res.bytes_stream(); + while let Some(chunk) = stream.next().await { + let chunk = chunk.map_err(|e| format!("stream read failed: {e}"))?; + buf.push_str(&String::from_utf8_lossy(&chunk)); + while let Some(nl) = buf.find('\n') { + let line: String = buf.drain(..=nl).collect(); + match parse(line.trim_end())? { + Some(StreamLine { token: Some(t), done }) => { + full.push_str(&t); + let _ = channel.send(LlmEvent::Token(t)); + if done { + return Ok(full); + } + } + Some(StreamLine { token: None, done: true }) => return Ok(full), + _ => {} + } + } + } + // Stream closed without an explicit done — return what we got. + Ok(full) +} + +async fn call_ollama_stream( + client: &reqwest::Client, + config: &crate::llm::LlmConfig, + messages: &[ChatMessage], + temperature: f32, + max_tokens: u32, + channel: &Channel, +) -> Result { + let url = format!("{}/api/chat", config.api_url.trim_end_matches('/')); + let body = json!({ + "model": config.model, + "messages": messages, + "options": { "temperature": temperature, "num_predict": max_tokens }, + }); + stream_chat(client, &url, body, channel, parse_ollama_line, None).await +} + +async fn call_openai_stream( + client: &reqwest::Client, + config: &crate::llm::LlmConfig, + messages: &[ChatMessage], + temperature: f32, + max_tokens: u32, + channel: &Channel, +) -> Result { + let url = format!("{}/v1/chat/completions", config.api_url.trim_end_matches('/')); + let body = json!({ + "model": config.model, + "messages": messages, + "temperature": temperature, + "max_tokens": max_tokens, + }); + let bearer = (!config.api_key.is_empty()).then_some(config.api_key.as_str()); + stream_chat(client, &url, body, channel, parse_openai_line, bearer).await +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn ollama_line_parses_tokens_and_done() { + let tok = parse_ollama_line("{\"message\":{\"role\":\"assistant\",\"content\":\"Hel\"},\"done\":false}").unwrap().unwrap(); + assert_eq!(tok.token.as_deref(), Some("Hel")); + assert!(!tok.done); + // final line: no message content, done + metrics + let fin = parse_ollama_line("{\"done\":true,\"total_duration\":123}").unwrap().unwrap(); + assert_eq!(fin.token, None); + assert!(fin.done); + // blank lines are skipped, errors bubble + assert!(parse_ollama_line("").unwrap().is_none()); + assert!(parse_ollama_line("{\"error\":\"model not found\"}").is_err()); + } + + #[test] + fn openai_line_parses_sse() { + let tok = parse_openai_line("data: {\"choices\":[{\"delta\":{\"content\":\"lo\"}}]}").unwrap().unwrap(); + assert_eq!(tok.token.as_deref(), Some("lo")); + let fin = parse_openai_line("data: [DONE]").unwrap().unwrap(); + assert!(fin.done); + // non-data lines (keep-alives) skipped; role-only deltas yield no token + assert!(parse_openai_line(": ping").unwrap().is_none()); + let role = parse_openai_line("data: {\"choices\":[{\"delta\":{\"role\":\"assistant\"}}]}").unwrap().unwrap(); + assert_eq!(role.token, None); + } } \ No newline at end of file