use futures_util::StreamExt;
use serde::{Deserialize, Serialize};
use tokio::sync::mpsc;
use tokio_util::sync::CancellationToken;
use crate::error::{LlmError, Result, SkadooshError};
use crate::llm::splitter::ClauseSplitter;
pub(crate) const CLAUSE_MIN_LEN: usize = 4;
pub(crate) const CLAUSE_MAX_LEN: usize = 160;
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct Message {
pub role: String,
pub content: String,
}
pub struct LlmClient {
http: reqwest::Client,
base_url: String,
model: String,
max_history_turns: usize,
history: Vec<Message>,
}
impl LlmClient {
pub fn new(base_url: &str, model: &str, system_prompt: &str, max_history_turns: usize) -> Self {
Self {
http: reqwest::Client::new(),
base_url: base_url.trim_end_matches('/').to_string(),
model: model.to_string(),
max_history_turns,
history: vec![Message {
role: "system".to_string(),
content: system_prompt.to_string(),
}],
}
}
pub async fn stream_reply(
&mut self,
user: &str,
clauses: mpsc::Sender<String>,
cancel: CancellationToken,
) -> Result<()> {
self.history.push(Message {
role: "user".to_string(),
content: user.to_string(),
});
let result = self.stream_reply_inner(&clauses, &cancel).await;
if matches!(result, Err(SkadooshError::Llm(LlmError::Cancelled))) {
self.truncate_history();
}
result
}
async fn stream_reply_inner(
&mut self,
clauses: &mpsc::Sender<String>,
cancel: &CancellationToken,
) -> Result<()> {
let body = serde_json::json!({
"model": self.model,
"messages": self.history,
"stream": true,
});
let url = format!("{}/chat/completions", self.base_url);
let resp = tokio::select! {
_ = cancel.cancelled() => return Err(LlmError::Cancelled.into()),
r = self.http.post(&url).json(&body).send() => r.map_err(LlmError::Http)?,
};
let resp = ensure_success(resp).await?;
let mut stream = resp.bytes_stream();
let mut splitter = ClauseSplitter::new(CLAUSE_MIN_LEN, CLAUSE_MAX_LEN);
let mut reply = String::new();
let mut lines = SseLineBuffer::default();
let mut done = false;
let mut eof = false;
while !done && !eof {
let chunk = tokio::select! {
_ = cancel.cancelled() => return Err(LlmError::Cancelled.into()),
c = stream.next() => c,
};
match chunk {
Some(Ok(bytes)) => lines.feed(&bytes),
Some(Err(e)) => return Err(LlmError::Http(e).into()),
None => {
lines.close();
eof = true;
}
}
while let Some(line) = lines.next_line() {
match parse_sse_line(&line) {
None => {}
Some(Ok(None)) => {
done = true;
break;
}
Some(Ok(Some(token))) => {
reply.push_str(&token);
for clause in splitter.push(&token) {
if !send_clause(clauses, cancel, clause).await? {
tracing::debug!("clauses receiver dropped; aborting LLM stream");
return Ok(());
}
}
}
Some(Err(e)) => {
tracing::warn!(error = %e, "skipping malformed SSE data line");
}
}
}
}
if let Some(rest) = splitter.flush() {
if !send_clause(clauses, cancel, rest).await? {
tracing::debug!("clauses receiver dropped at stream end");
return Ok(());
}
}
self.history.push(Message {
role: "assistant".to_string(),
content: reply,
});
self.truncate_history();
Ok(())
}
pub fn history(&self) -> &[Message] {
&self.history
}
fn truncate_history(&mut self) {
let keep = 2 * self.max_history_turns;
if self.history.len() > 1 + keep {
let drop = self.history.len() - 1 - keep;
self.history.drain(1..=drop);
}
}
}
pub(crate) async fn ensure_success(resp: reqwest::Response) -> Result<reqwest::Response> {
let status = resp.status();
if status.is_success() {
return Ok(resp);
}
let text = resp.text().await.unwrap_or_default();
Err(LlmError::Api {
status: status.as_u16(),
body: text.chars().take(1024).collect(),
}
.into())
}
async fn send_clause(
clauses: &mpsc::Sender<String>,
cancel: &CancellationToken,
clause: String,
) -> std::result::Result<bool, LlmError> {
tokio::select! {
_ = cancel.cancelled() => Err(LlmError::Cancelled),
sent = clauses.send(clause) => Ok(sent.is_ok()),
}
}
#[derive(Default)]
pub(crate) struct SseLineBuffer {
buf: Vec<u8>,
eof: bool,
}
impl SseLineBuffer {
pub(crate) fn feed(&mut self, chunk: &[u8]) {
self.buf.extend_from_slice(chunk);
}
pub(crate) fn close(&mut self) {
self.eof = true;
}
pub(crate) fn next_line(&mut self) -> Option<String> {
if let Some(nl) = self.buf.iter().position(|&b| b == b'\n') {
let line_bytes: Vec<u8> = self.buf.drain(..=nl).collect();
return Some(String::from_utf8_lossy(&line_bytes).into_owned());
}
if self.eof && !self.buf.is_empty() {
let rest = std::mem::take(&mut self.buf);
return Some(String::from_utf8_lossy(&rest).into_owned());
}
None
}
}
pub fn parse_sse_line(line: &str) -> Option<Result<Option<String>>> {
let line = line.trim();
if line.is_empty() || line.starts_with(':') {
return None;
}
let data = line.strip_prefix("data:")?;
let data = data.trim();
if data == "[DONE]" {
return Some(Ok(None));
}
let parsed: serde_json::Value = match serde_json::from_str(data) {
Ok(v) => v,
Err(e) => {
return Some(Err(LlmError::Sse(format!("malformed SSE data: {e}")).into()));
}
};
let token = parsed
.get("choices")?
.as_array()?
.first()?
.get("delta")?
.get("content")?
.as_str()?;
if token.is_empty() {
None
} else {
Some(Ok(Some(token.to_string())))
}
}