feat: real token streaming from Ollama (NDJSON) + OpenAI-compatible (SSE)
generate_stream previously did a non-streaming call and emitted the whole response as one fake token. Now reqwest bytes_stream feeds a line buffer; every token piece goes out as it arrives via the existing Channel. SessionLogger's progressive display lights up unchanged. Per-line parsers are pure fns with unit tests; unparseable lines are skipped, server error lines bubble.
This commit is contained in:
@@ -63,17 +63,16 @@ pub async fn generate_stream(
|
|||||||
|
|
||||||
emit_busy(&app, true);
|
emit_busy(&app, true);
|
||||||
tauri::async_runtime::spawn(async move {
|
tauri::async_runtime::spawn(async move {
|
||||||
// For now, we do a non-streaming call and emit the full response as one token
|
// Real token streaming: Ollama NDJSON / OpenAI SSE, one LlmEvent::Token
|
||||||
// Real SSE streaming from Ollama/OpenAI can be added later
|
// per piece as it arrives. Falls back to Done carrying the full text.
|
||||||
let result = if llm::is_ollama(&config.provider, &config.api_url) {
|
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 {
|
} else {
|
||||||
call_openai(&client, &config, &messages, temperature, max_tokens).await
|
call_openai_stream(&client, &config, &messages, temperature, max_tokens, &channel).await
|
||||||
};
|
};
|
||||||
|
|
||||||
match result {
|
match result {
|
||||||
Ok(text) => {
|
Ok(text) => {
|
||||||
let _ = channel.send(LlmEvent::Token(text.clone()));
|
|
||||||
let _ = channel.send(LlmEvent::Done(text));
|
let _ = channel.send(LlmEvent::Done(text));
|
||||||
}
|
}
|
||||||
Err(e) => {
|
Err(e) => {
|
||||||
@@ -305,3 +304,176 @@ async fn call_openai(
|
|||||||
.map(|c| c.message.content.clone())
|
.map(|c| c.message.content.clone())
|
||||||
.ok_or_else(|| "No response from API".to_string())
|
.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<String>,
|
||||||
|
done: bool,
|
||||||
|
}
|
||||||
|
|
||||||
|
/// Parse one Ollama NDJSON line: `{"message":{"content":"…"},"done":false}`.
|
||||||
|
fn parse_ollama_line(line: &str) -> Result<Option<StreamLine>, 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<Option<StreamLine>, 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<LlmEvent>,
|
||||||
|
parse: fn(&str) -> Result<Option<StreamLine>, String>,
|
||||||
|
bearer: Option<&str>,
|
||||||
|
) -> Result<String, String> {
|
||||||
|
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<LlmEvent>,
|
||||||
|
parse: fn(&str) -> Result<Option<StreamLine>, String>,
|
||||||
|
) -> Result<String, String> {
|
||||||
|
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<LlmEvent>,
|
||||||
|
) -> Result<String, String> {
|
||||||
|
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<LlmEvent>,
|
||||||
|
) -> Result<String, String> {
|
||||||
|
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);
|
||||||
|
}
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user