feat: chat typing

This commit is contained in:
2026-04-10 21:33:52 +02:00
parent d5856557b4
commit e31cdf131f
6 changed files with 277 additions and 221 deletions
+129 -152
View File
@@ -1,10 +1,9 @@
use crate::dto::{api, ollama};
use crate::errors::OllamaError;
use axum::Json;
use axum::response::sse::Event;
use futures::StreamExt;
use reqwest::Client;
use serde_json::{Value, json};
use serde_json::json;
use tokio_stream::wrappers::ReceiverStream;
#[derive(Clone)]
@@ -48,80 +47,14 @@ impl OllamaProvider {
Ok((prompt, model))
}
// fn format_completion_response(&self, res: &Value) -> Value {
// json!({
// "id": "cmpl-ollama",
// "object": "text_completion",
// "model": res.get("model"),
// "choices": [{
// "text": res.get("response"),
// "index": 0,
// "finish_reason": if res.get("done").and_then(|v| v.as_bool()).unwrap_or(false) {
// "stop"
// } else {
// "length"
// },
// }],
// "usage": {
// "prompt_tokens": res.get("prompt_eval_count"),
// "completion_tokens": res.get("eval_count"),
// "total_tokens": null,
// }
// })
// }
fn extract_chat_params<'a>(
&self,
body: &'a api::ChatRequest,
) -> Result<(&'a [api::Message], &'a str), OllamaError> {
let model = body.base.model.as_str();
// fn extract_chat_params<'a>(
// &self,
// body: &'a Value,
// ) -> Result<(&'a str, &'a Vec<Value>), OllamaError> {
// let model = body
// .get("model")
// .and_then(|v| v.as_str())
// .unwrap_or("llama3");
// let messages = body
// .get("messages")
// .and_then(|v| v.as_array())
// .filter(|arr| !arr.is_empty())
// .ok_or(OllamaError::MissingMessages)?;
// let has_user_msg = messages
// .iter()
// .any(|m| m.get("role").and_then(|r| r.as_str()) == Some("user"));
// if !has_user_msg {
// return Err(OllamaError::MissingMessages);
// }
// Ok((model, messages))
// }
// fn format_chat_response(&self, res: &Value) -> Value {
// json!({
// "id": "chatcmpl-ollama",
// "object": "chat.completion",
// "model": res.get("model"),
// "choices": [{
// "index": 0,
// "message": {
// "role": res.get("message").and_then(|m| m.get("role")),
// "content": res.get("message").and_then(|m| m.get("content")),
// },
// "finish_reason": res
// .get("done_reason")
// .and_then(|v| v.as_str())
// .unwrap_or("stop"),
// }],
// "usage": {
// "prompt_tokens": res.get("prompt_eval_count"),
// "completion_tokens": res.get("eval_count"),
// "total_tokens": res.get("prompt_eval_count")
// .and_then(|p| p.as_u64())
// .zip(res.get("eval_count").and_then(|e| e.as_u64()))
// .map(|(p, e)| p + e),
// }
// })
// }
Ok((&body.messages, model))
}
// // ── public endpoints ─────────────────────────────────────────────────────
@@ -166,7 +99,7 @@ impl OllamaProvider {
.json(&payload)
.send()
.await?
.json::<Value>()
.json::<ollama::OllamaGenerateResponse>()
.await?;
Ok(api::LoadModelResponse {
@@ -197,7 +130,7 @@ impl OllamaProvider {
.json(&payload)
.send()
.await?
.json::<Value>()
.json::<ollama::OllamaGenerateResponse>()
.await?;
Ok(api::UnloadModelResponse {
@@ -227,7 +160,7 @@ impl OllamaProvider {
return Err(OllamaError::ModelNotFound(model.to_string()));
}
let options = ollama::OllamaOptions::from(body);
let options = ollama::OllamaOptions::from(&body.base);
let payload = ollama::OllamaGenerateRequest {
model,
@@ -269,7 +202,7 @@ impl OllamaProvider {
return Err(OllamaError::ModelNotFound(model.to_string()));
}
let options = ollama::OllamaOptions::from(body);
let options = ollama::OllamaOptions::from(&body.base);
let payload = ollama::OllamaGenerateRequest {
model,
@@ -332,90 +265,134 @@ impl OllamaProvider {
Ok(ReceiverStream::new(rx))
}
// pub async fn chat_completions(&self, body: Value) -> Result<Value, OllamaError> {
// let (model, messages) = self.extract_chat_params(&body)?;
// self.validate_model(model).await?;
pub async fn chat_completions(
&self,
body: &api::ChatRequest,
) -> Result<api::ChatCompletionResponse, OllamaError> {
let url = format!("{}/api/chat", self.base_url);
// let payload = json!({
// "model": model,
// "messages": messages,
// "stream": false,
// "options": Self::build_options(&body),
// });
let (messages, model) = self.extract_chat_params(body)?;
// let res = self
// .client
// .post(format!("{}/api/chat", self.base_url))
// .json(&payload)
// .send()
// .await?
// .json::<Value>()
// .await?;
if body.messages.is_empty() {
return Err(OllamaError::MissingMessages);
}
// Ok(self.format_chat_response(&res))
// }
if model.is_empty() {
return Err(OllamaError::MissingModel);
}
// pub async fn chat_completions_stream(
// &self,
// body: Value,
// ) -> Result<ReceiverStream<Result<Event, OllamaError>>, OllamaError> {
// let (model, messages) = self.extract_chat_params(&body)?;
// self.validate_model(model).await?;
let exists = self.model_exists(model).await?;
if !exists {
return Err(OllamaError::ModelNotFound(model.to_string()));
}
// let payload = json!({
// "model": model,
// "messages": messages,
// "stream": true,
// "options": Self::build_options(&body),
// });
// let options = ollama::OllamaOptions::from(body);
let options = ollama::OllamaOptions::from(&body.base);
// let mut byte_stream = self
// .client
// .post(format!("{}/api/chat", self.base_url))
// .json(&payload)
// .send()
// .await?
// .bytes_stream();
let payload = ollama::OllamaChatRequest {
model,
messages,
stream: false,
options,
};
// let (tx, rx) = tokio::sync::mpsc::channel(32);
let res = self
.client
.post(url)
.json(&payload)
.send()
.await?
.json::<ollama::OllamaChatResponse>()
.await?;
// tokio::spawn(async move {
// while let Some(chunk) = byte_stream.next().await {
// let chunk = match chunk {
// Ok(b) => b,
// Err(e) => {
// let _ = tx.send(Err(OllamaError::Http(e))).await;
// break;
// }
// };
Ok(api::ChatCompletionResponse::from(res))
}
// if let Ok(json) = serde_json::from_slice::<Value>(&chunk) {
// let done = json.get("done").and_then(|v| v.as_bool()).unwrap_or(false);
pub async fn chat_completions_stream(
&self,
body: &api::ChatRequest,
) -> Result<ReceiverStream<Result<Event, OllamaError>>, OllamaError> {
let url = format!("{}/api/chat", self.base_url);
// let event_data = serde_json::to_string(&json!({
// "id": "chatcmpl-ollama",
// "object": "chat.completion.chunk",
// "choices": [{
// "index": 0,
// "delta": {
// "role": json.get("message").and_then(|m| m.get("role")),
// "content": json.get("message").and_then(|m| m.get("content")),
// },
// "finish_reason": if done { json!("stop") } else { json!(null) },
// }],
// }))
// .unwrap_or_default();
let (messages, model) = self.extract_chat_params(body)?;
// let _ = tx.send(Ok(Event::default().data(event_data))).await;
if messages.is_empty() {
return Err(OllamaError::MissingMessages);
}
// if done {
// let _ = tx.send(Ok(Event::default().data("[DONE]"))).await;
// break;
// }
// }
// }
// });
if model.is_empty() {
return Err(OllamaError::MissingModel);
}
// Ok(ReceiverStream::new(rx))
// }
let exists = self.model_exists(model).await?;
if !exists {
return Err(OllamaError::ModelNotFound(model.to_string()));
}
let options = ollama::OllamaOptions::from(&body.base);
let payload = ollama::OllamaChatRequest {
model,
messages,
stream: true,
options,
};
let mut byte_stream = self
.client
.post(url)
.json(&payload)
.send()
.await?
.bytes_stream();
let (tx, rx) = tokio::sync::mpsc::channel(32);
tokio::spawn(async move {
let stream_id = format!("chatcmpl-{}", uuid::Uuid::new_v4());
while let Some(chunk) = byte_stream.next().await {
let chunk = match chunk {
Ok(b) => b,
Err(e) => {
let _ = tx.send(Err(OllamaError::Http(e))).await;
break;
}
};
let parsed: ollama::OllamaChatResponse = match serde_json::from_slice(&chunk) {
Ok(v) => v,
Err(_) => continue,
};
let event = api::ChatCompletionChunk {
id: stream_id.clone(),
object: "chat.completion.chunk".to_string(),
choices: vec![api::ChatChunkChoice {
index: 0,
delta: api::ChatDelta {
role: Some(parsed.message.role),
content: Some(parsed.message.content),
},
finish_reason: if parsed.done {
Some(api::FinishReason::Stop)
} else {
None
},
}],
};
let event_data = serde_json::to_string(&event).unwrap_or_default();
let _ = tx.send(Ok(Event::default().data(event_data))).await;
if parsed.done {
let _ = tx.send(Ok(Event::default().data("[DONE]"))).await;
break;
}
}
});
Ok(ReceiverStream::new(rx))
}
}
+44 -14
View File
@@ -18,20 +18,6 @@ impl From<ollama::OllamaModel> for api::ModelInfo {
}
}
impl From<&api::CompletionRequest> for ollama::OllamaOptions {
fn from(req: &api::CompletionRequest) -> Self {
Self {
temperature: req.base.temperature,
top_p: req.base.top_p,
top_k: req.base.top_k,
repeat_penalty: req.base.repeat_penalty,
seed: req.base.seed,
num_ctx: req.base.num_ctx,
num_predict: req.base.num_predict,
}
}
}
impl From<ollama::OllamaGenerateResponse> for api::CompletionResponse {
fn from(res: ollama::OllamaGenerateResponse) -> Self {
Self {
@@ -54,3 +40,47 @@ impl From<ollama::OllamaGenerateResponse> for api::CompletionResponse {
}
}
}
impl From<&api::BaseLLMRequest> for ollama::OllamaOptions {
fn from(base: &api::BaseLLMRequest) -> Self {
Self {
temperature: base.temperature,
top_p: base.top_p,
top_k: base.top_k,
repeat_penalty: base.repeat_penalty,
seed: base.seed,
num_ctx: base.num_ctx,
num_predict: base.num_predict,
}
}
}
impl From<ollama::OllamaChatResponse> for api::ChatCompletionResponse {
fn from(res: ollama::OllamaChatResponse) -> Self {
let prompt_tokens = res.prompt_eval_count.unwrap_or(0);
let completion_tokens = res.eval_count.unwrap_or(0);
Self {
id: Uuid::new_v4().to_string(),
object: "chat.completion".to_string(),
created: Utc::now().timestamp() as u64,
model: res.model,
choices: vec![api::ChatChoice {
index: 0,
message: res.message,
finish_reason: if res.done {
api::FinishReason::Stop
} else {
api::FinishReason::Length
},
}],
usage: Some(api::Usage {
prompt_tokens,
completion_tokens,
total_tokens: prompt_tokens + completion_tokens,
}),
}
}
}