feat: chat typing
This commit is contained in:
+59
-21
@@ -3,27 +3,6 @@ use utoipa::ToSchema;
|
|||||||
|
|
||||||
use crate::errors::OllamaError;
|
use crate::errors::OllamaError;
|
||||||
|
|
||||||
#[derive(Debug, Serialize, Deserialize, ToSchema)]
|
|
||||||
#[serde(rename_all = "lowercase")]
|
|
||||||
pub enum Role {
|
|
||||||
System,
|
|
||||||
User,
|
|
||||||
Assistant,
|
|
||||||
}
|
|
||||||
|
|
||||||
#[derive(Serialize, Deserialize, ToSchema)]
|
|
||||||
pub struct ChatResponse {
|
|
||||||
pub response: String,
|
|
||||||
}
|
|
||||||
|
|
||||||
#[derive(Debug, Serialize, Deserialize, ToSchema)]
|
|
||||||
pub struct Message {
|
|
||||||
pub role: Role,
|
|
||||||
pub content: String,
|
|
||||||
}
|
|
||||||
|
|
||||||
// ---------------------------
|
|
||||||
|
|
||||||
#[derive(Debug, Serialize, Deserialize, ToSchema)]
|
#[derive(Debug, Serialize, Deserialize, ToSchema)]
|
||||||
pub struct ModelsResponse {
|
pub struct ModelsResponse {
|
||||||
pub models: Vec<ModelInfo>,
|
pub models: Vec<ModelInfo>,
|
||||||
@@ -155,3 +134,62 @@ pub struct CompletionChunk {
|
|||||||
pub object: String,
|
pub object: String,
|
||||||
pub choices: Vec<Choice>,
|
pub choices: Vec<Choice>,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[derive(Debug, Deserialize, Serialize, ToSchema)]
|
||||||
|
pub struct ChatRequest {
|
||||||
|
#[serde(flatten)]
|
||||||
|
pub base: BaseLLMRequest,
|
||||||
|
|
||||||
|
pub messages: Vec<Message>,
|
||||||
|
}
|
||||||
|
|
||||||
|
#[derive(Debug, Deserialize, Serialize, ToSchema)]
|
||||||
|
pub struct Message {
|
||||||
|
pub role: Role,
|
||||||
|
pub content: String,
|
||||||
|
}
|
||||||
|
|
||||||
|
#[derive(Debug, Deserialize, Serialize, ToSchema)]
|
||||||
|
#[serde(rename_all = "lowercase")]
|
||||||
|
pub enum Role {
|
||||||
|
System,
|
||||||
|
User,
|
||||||
|
Assistant,
|
||||||
|
}
|
||||||
|
|
||||||
|
#[derive(Debug, Serialize, Deserialize, ToSchema)]
|
||||||
|
pub struct ChatCompletionResponse {
|
||||||
|
pub id: String,
|
||||||
|
pub object: String,
|
||||||
|
pub created: u64,
|
||||||
|
pub model: String,
|
||||||
|
pub choices: Vec<ChatChoice>,
|
||||||
|
pub usage: Option<Usage>, // optional (Ollama may not always provide)
|
||||||
|
}
|
||||||
|
|
||||||
|
#[derive(Debug, Serialize, Deserialize, ToSchema)]
|
||||||
|
pub struct ChatChoice {
|
||||||
|
pub index: u32,
|
||||||
|
pub message: Message,
|
||||||
|
pub finish_reason: FinishReason,
|
||||||
|
}
|
||||||
|
|
||||||
|
#[derive(Debug, Serialize)]
|
||||||
|
pub struct ChatCompletionChunk {
|
||||||
|
pub id: String,
|
||||||
|
pub object: String,
|
||||||
|
pub choices: Vec<ChatChunkChoice>,
|
||||||
|
}
|
||||||
|
|
||||||
|
#[derive(Debug, Serialize)]
|
||||||
|
pub struct ChatChunkChoice {
|
||||||
|
pub index: u32,
|
||||||
|
pub delta: ChatDelta,
|
||||||
|
pub finish_reason: Option<FinishReason>,
|
||||||
|
}
|
||||||
|
|
||||||
|
#[derive(Debug, Serialize)]
|
||||||
|
pub struct ChatDelta {
|
||||||
|
pub role: Option<Role>,
|
||||||
|
pub content: Option<String>,
|
||||||
|
}
|
||||||
|
|||||||
+20
-1
@@ -1,5 +1,6 @@
|
|||||||
use serde::{Deserialize, Serialize};
|
use serde::{Deserialize, Serialize};
|
||||||
use utoipa::ToSchema;
|
|
||||||
|
use crate::dto::api;
|
||||||
|
|
||||||
#[derive(Debug, Serialize, Deserialize)]
|
#[derive(Debug, Serialize, Deserialize)]
|
||||||
pub struct OllamaModels {
|
pub struct OllamaModels {
|
||||||
@@ -59,3 +60,21 @@ pub struct OllamaGenerateResponse {
|
|||||||
pub prompt_eval_count: Option<u32>,
|
pub prompt_eval_count: Option<u32>,
|
||||||
pub eval_count: Option<u32>,
|
pub eval_count: Option<u32>,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[derive(Debug, Serialize)]
|
||||||
|
pub struct OllamaChatRequest<'a> {
|
||||||
|
pub model: &'a str,
|
||||||
|
pub messages: &'a [api::Message],
|
||||||
|
pub stream: bool,
|
||||||
|
pub options: OllamaOptions,
|
||||||
|
}
|
||||||
|
|
||||||
|
#[derive(Debug, Deserialize)]
|
||||||
|
pub struct OllamaChatResponse {
|
||||||
|
pub model: String,
|
||||||
|
pub message: api::Message,
|
||||||
|
pub done: bool,
|
||||||
|
|
||||||
|
pub prompt_eval_count: Option<u32>,
|
||||||
|
pub eval_count: Option<u32>,
|
||||||
|
}
|
||||||
|
|||||||
+129
-152
@@ -1,10 +1,9 @@
|
|||||||
use crate::dto::{api, ollama};
|
use crate::dto::{api, ollama};
|
||||||
use crate::errors::OllamaError;
|
use crate::errors::OllamaError;
|
||||||
use axum::Json;
|
|
||||||
use axum::response::sse::Event;
|
use axum::response::sse::Event;
|
||||||
use futures::StreamExt;
|
use futures::StreamExt;
|
||||||
use reqwest::Client;
|
use reqwest::Client;
|
||||||
use serde_json::{Value, json};
|
use serde_json::json;
|
||||||
use tokio_stream::wrappers::ReceiverStream;
|
use tokio_stream::wrappers::ReceiverStream;
|
||||||
|
|
||||||
#[derive(Clone)]
|
#[derive(Clone)]
|
||||||
@@ -48,80 +47,14 @@ impl OllamaProvider {
|
|||||||
Ok((prompt, model))
|
Ok((prompt, model))
|
||||||
}
|
}
|
||||||
|
|
||||||
// fn format_completion_response(&self, res: &Value) -> Value {
|
fn extract_chat_params<'a>(
|
||||||
// json!({
|
&self,
|
||||||
// "id": "cmpl-ollama",
|
body: &'a api::ChatRequest,
|
||||||
// "object": "text_completion",
|
) -> Result<(&'a [api::Message], &'a str), OllamaError> {
|
||||||
// "model": res.get("model"),
|
let model = body.base.model.as_str();
|
||||||
// "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>(
|
Ok((&body.messages, model))
|
||||||
// &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),
|
|
||||||
// }
|
|
||||||
// })
|
|
||||||
// }
|
|
||||||
|
|
||||||
// // ── public endpoints ─────────────────────────────────────────────────────
|
// // ── public endpoints ─────────────────────────────────────────────────────
|
||||||
|
|
||||||
@@ -166,7 +99,7 @@ impl OllamaProvider {
|
|||||||
.json(&payload)
|
.json(&payload)
|
||||||
.send()
|
.send()
|
||||||
.await?
|
.await?
|
||||||
.json::<Value>()
|
.json::<ollama::OllamaGenerateResponse>()
|
||||||
.await?;
|
.await?;
|
||||||
|
|
||||||
Ok(api::LoadModelResponse {
|
Ok(api::LoadModelResponse {
|
||||||
@@ -197,7 +130,7 @@ impl OllamaProvider {
|
|||||||
.json(&payload)
|
.json(&payload)
|
||||||
.send()
|
.send()
|
||||||
.await?
|
.await?
|
||||||
.json::<Value>()
|
.json::<ollama::OllamaGenerateResponse>()
|
||||||
.await?;
|
.await?;
|
||||||
|
|
||||||
Ok(api::UnloadModelResponse {
|
Ok(api::UnloadModelResponse {
|
||||||
@@ -227,7 +160,7 @@ impl OllamaProvider {
|
|||||||
return Err(OllamaError::ModelNotFound(model.to_string()));
|
return Err(OllamaError::ModelNotFound(model.to_string()));
|
||||||
}
|
}
|
||||||
|
|
||||||
let options = ollama::OllamaOptions::from(body);
|
let options = ollama::OllamaOptions::from(&body.base);
|
||||||
|
|
||||||
let payload = ollama::OllamaGenerateRequest {
|
let payload = ollama::OllamaGenerateRequest {
|
||||||
model,
|
model,
|
||||||
@@ -269,7 +202,7 @@ impl OllamaProvider {
|
|||||||
return Err(OllamaError::ModelNotFound(model.to_string()));
|
return Err(OllamaError::ModelNotFound(model.to_string()));
|
||||||
}
|
}
|
||||||
|
|
||||||
let options = ollama::OllamaOptions::from(body);
|
let options = ollama::OllamaOptions::from(&body.base);
|
||||||
|
|
||||||
let payload = ollama::OllamaGenerateRequest {
|
let payload = ollama::OllamaGenerateRequest {
|
||||||
model,
|
model,
|
||||||
@@ -332,90 +265,134 @@ impl OllamaProvider {
|
|||||||
Ok(ReceiverStream::new(rx))
|
Ok(ReceiverStream::new(rx))
|
||||||
}
|
}
|
||||||
|
|
||||||
// pub async fn chat_completions(&self, body: Value) -> Result<Value, OllamaError> {
|
pub async fn chat_completions(
|
||||||
// let (model, messages) = self.extract_chat_params(&body)?;
|
&self,
|
||||||
// self.validate_model(model).await?;
|
body: &api::ChatRequest,
|
||||||
|
) -> Result<api::ChatCompletionResponse, OllamaError> {
|
||||||
|
let url = format!("{}/api/chat", self.base_url);
|
||||||
|
|
||||||
// let payload = json!({
|
let (messages, model) = self.extract_chat_params(body)?;
|
||||||
// "model": model,
|
|
||||||
// "messages": messages,
|
|
||||||
// "stream": false,
|
|
||||||
// "options": Self::build_options(&body),
|
|
||||||
// });
|
|
||||||
|
|
||||||
// let res = self
|
if body.messages.is_empty() {
|
||||||
// .client
|
return Err(OllamaError::MissingMessages);
|
||||||
// .post(format!("{}/api/chat", self.base_url))
|
}
|
||||||
// .json(&payload)
|
|
||||||
// .send()
|
|
||||||
// .await?
|
|
||||||
// .json::<Value>()
|
|
||||||
// .await?;
|
|
||||||
|
|
||||||
// Ok(self.format_chat_response(&res))
|
if model.is_empty() {
|
||||||
// }
|
return Err(OllamaError::MissingModel);
|
||||||
|
}
|
||||||
|
|
||||||
// pub async fn chat_completions_stream(
|
let exists = self.model_exists(model).await?;
|
||||||
// &self,
|
if !exists {
|
||||||
// body: Value,
|
return Err(OllamaError::ModelNotFound(model.to_string()));
|
||||||
// ) -> Result<ReceiverStream<Result<Event, OllamaError>>, OllamaError> {
|
}
|
||||||
// let (model, messages) = self.extract_chat_params(&body)?;
|
|
||||||
// self.validate_model(model).await?;
|
|
||||||
|
|
||||||
// let payload = json!({
|
// let options = ollama::OllamaOptions::from(body);
|
||||||
// "model": model,
|
let options = ollama::OllamaOptions::from(&body.base);
|
||||||
// "messages": messages,
|
|
||||||
// "stream": true,
|
|
||||||
// "options": Self::build_options(&body),
|
|
||||||
// });
|
|
||||||
|
|
||||||
// let mut byte_stream = self
|
let payload = ollama::OllamaChatRequest {
|
||||||
// .client
|
model,
|
||||||
// .post(format!("{}/api/chat", self.base_url))
|
messages,
|
||||||
// .json(&payload)
|
stream: false,
|
||||||
// .send()
|
options,
|
||||||
// .await?
|
};
|
||||||
// .bytes_stream();
|
|
||||||
|
|
||||||
// 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 {
|
Ok(api::ChatCompletionResponse::from(res))
|
||||||
// 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;
|
|
||||||
// }
|
|
||||||
// };
|
|
||||||
|
|
||||||
// if let Ok(json) = serde_json::from_slice::<Value>(&chunk) {
|
pub async fn chat_completions_stream(
|
||||||
// let done = json.get("done").and_then(|v| v.as_bool()).unwrap_or(false);
|
&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!({
|
let (messages, model) = self.extract_chat_params(body)?;
|
||||||
// "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 _ = tx.send(Ok(Event::default().data(event_data))).await;
|
if messages.is_empty() {
|
||||||
|
return Err(OllamaError::MissingMessages);
|
||||||
|
}
|
||||||
|
|
||||||
// if done {
|
if model.is_empty() {
|
||||||
// let _ = tx.send(Ok(Event::default().data("[DONE]"))).await;
|
return Err(OllamaError::MissingModel);
|
||||||
// break;
|
}
|
||||||
// }
|
|
||||||
// }
|
|
||||||
// }
|
|
||||||
// });
|
|
||||||
|
|
||||||
// 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))
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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 {
|
impl From<ollama::OllamaGenerateResponse> for api::CompletionResponse {
|
||||||
fn from(res: ollama::OllamaGenerateResponse) -> Self {
|
fn from(res: ollama::OllamaGenerateResponse) -> Self {
|
||||||
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,
|
||||||
|
}),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
+23
-31
@@ -6,7 +6,6 @@ use axum::{
|
|||||||
sse::{KeepAlive, Sse},
|
sse::{KeepAlive, Sse},
|
||||||
},
|
},
|
||||||
};
|
};
|
||||||
use serde_json::Value;
|
|
||||||
|
|
||||||
use crate::dto::api;
|
use crate::dto::api;
|
||||||
use crate::errors::OllamaError;
|
use crate::errors::OllamaError;
|
||||||
@@ -23,9 +22,7 @@ pub async fn completions(
|
|||||||
State(state): State<AppState>,
|
State(state): State<AppState>,
|
||||||
Json(body): Json<api::CompletionRequest>,
|
Json(body): Json<api::CompletionRequest>,
|
||||||
) -> Result<Response, (axum::http::StatusCode, String)> {
|
) -> Result<Response, (axum::http::StatusCode, String)> {
|
||||||
let wants_stream = body.base.stream;
|
if body.base.stream {
|
||||||
|
|
||||||
if wants_stream {
|
|
||||||
let stream = state
|
let stream = state
|
||||||
.ollama
|
.ollama
|
||||||
.completions_stream(&body)
|
.completions_stream(&body)
|
||||||
@@ -42,35 +39,30 @@ pub async fn completions(
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// pub async fn chat_completions(
|
pub async fn chat_completions(
|
||||||
// State(state): State<AppState>,
|
State(state): State<AppState>,
|
||||||
// Json(body): Json<Value>,
|
Json(body): Json<api::ChatRequest>,
|
||||||
// ) -> Result<Response, (axum::http::StatusCode, String)> {
|
) -> Result<Response, (axum::http::StatusCode, String)> {
|
||||||
// let wants_stream = body
|
if body.base.stream {
|
||||||
// .get("stream")
|
let stream = state
|
||||||
// .and_then(|v| v.as_bool())
|
.ollama
|
||||||
// .unwrap_or(false);
|
.chat_completions_stream(&body)
|
||||||
|
.await
|
||||||
|
.map_err(ollama_err)?;
|
||||||
|
|
||||||
// if wants_stream {
|
Ok(Sse::new(stream)
|
||||||
// let stream = state
|
.keep_alive(KeepAlive::default())
|
||||||
// .ollama
|
.into_response())
|
||||||
// .chat_completions_stream(body)
|
} else {
|
||||||
// .await
|
let response = state
|
||||||
// .map_err(ollama_err)?;
|
.ollama
|
||||||
|
.chat_completions(&body)
|
||||||
|
.await
|
||||||
|
.map_err(ollama_err)?;
|
||||||
|
|
||||||
// Ok(Sse::new(stream)
|
Ok(Json(response).into_response())
|
||||||
// .keep_alive(KeepAlive::default())
|
}
|
||||||
// .into_response())
|
}
|
||||||
// } else {
|
|
||||||
// let response = state
|
|
||||||
// .ollama
|
|
||||||
// .chat_completions(body)
|
|
||||||
// .await
|
|
||||||
// .map_err(ollama_err)?;
|
|
||||||
|
|
||||||
// Ok(Json(response).into_response())
|
|
||||||
// }
|
|
||||||
// }
|
|
||||||
|
|
||||||
fn ollama_err(e: OllamaError) -> (axum::http::StatusCode, String) {
|
fn ollama_err(e: OllamaError) -> (axum::http::StatusCode, String) {
|
||||||
match e {
|
match e {
|
||||||
|
|||||||
@@ -14,7 +14,7 @@ pub fn protected_router() -> Router<AppState> {
|
|||||||
Router::new()
|
Router::new()
|
||||||
.route("/models", get(models::list_models))
|
.route("/models", get(models::list_models))
|
||||||
.route("/completions", post(chat::completions))
|
.route("/completions", post(chat::completions))
|
||||||
// .route("/chat/completions", post(chat::chat_completions))
|
.route("/chat/completions", post(chat::chat_completions))
|
||||||
.route("/models/{model}/load", post(models::load_model))
|
.route("/models/{model}/load", post(models::load_model))
|
||||||
.route("/models/{model}/unload", post(models::unload_model))
|
.route("/models/{model}/unload", post(models::unload_model))
|
||||||
}
|
}
|
||||||
@@ -22,5 +22,5 @@ pub fn protected_router() -> Router<AppState> {
|
|||||||
pub fn router() -> Router<AppState> {
|
pub fn router() -> Router<AppState> {
|
||||||
Router::new()
|
Router::new()
|
||||||
// .merge(public_router())
|
// .merge(public_router())
|
||||||
.merge(protected_router()) //.layer(middleware::from_fn(auth_middleware)))
|
.merge(protected_router().layer(middleware::from_fn(auth_middleware)))
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user