Generated
+29
@@ -168,6 +168,7 @@ version = "0.1.0"
|
|||||||
dependencies = [
|
dependencies = [
|
||||||
"axum",
|
"axum",
|
||||||
"dotenvy",
|
"dotenvy",
|
||||||
|
"futures",
|
||||||
"jsonwebtoken",
|
"jsonwebtoken",
|
||||||
"once_cell",
|
"once_cell",
|
||||||
"reqwest",
|
"reqwest",
|
||||||
@@ -175,6 +176,7 @@ dependencies = [
|
|||||||
"serde_json",
|
"serde_json",
|
||||||
"thiserror 2.0.18",
|
"thiserror 2.0.18",
|
||||||
"tokio",
|
"tokio",
|
||||||
|
"tokio-stream",
|
||||||
"wiremock",
|
"wiremock",
|
||||||
]
|
]
|
||||||
|
|
||||||
@@ -1156,6 +1158,7 @@ dependencies = [
|
|||||||
"bytes",
|
"bytes",
|
||||||
"encoding_rs",
|
"encoding_rs",
|
||||||
"futures-core",
|
"futures-core",
|
||||||
|
"futures-util",
|
||||||
"h2",
|
"h2",
|
||||||
"http",
|
"http",
|
||||||
"http-body",
|
"http-body",
|
||||||
@@ -1177,12 +1180,14 @@ dependencies = [
|
|||||||
"sync_wrapper",
|
"sync_wrapper",
|
||||||
"tokio",
|
"tokio",
|
||||||
"tokio-rustls",
|
"tokio-rustls",
|
||||||
|
"tokio-util",
|
||||||
"tower",
|
"tower",
|
||||||
"tower-http",
|
"tower-http",
|
||||||
"tower-service",
|
"tower-service",
|
||||||
"url",
|
"url",
|
||||||
"wasm-bindgen",
|
"wasm-bindgen",
|
||||||
"wasm-bindgen-futures",
|
"wasm-bindgen-futures",
|
||||||
|
"wasm-streams",
|
||||||
"web-sys",
|
"web-sys",
|
||||||
]
|
]
|
||||||
|
|
||||||
@@ -1663,6 +1668,17 @@ dependencies = [
|
|||||||
"tokio",
|
"tokio",
|
||||||
]
|
]
|
||||||
|
|
||||||
|
[[package]]
|
||||||
|
name = "tokio-stream"
|
||||||
|
version = "0.1.18"
|
||||||
|
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
|
checksum = "32da49809aab5c3bc678af03902d4ccddea2a87d028d86392a4b1560c6906c70"
|
||||||
|
dependencies = [
|
||||||
|
"futures-core",
|
||||||
|
"pin-project-lite",
|
||||||
|
"tokio",
|
||||||
|
]
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
name = "tokio-util"
|
name = "tokio-util"
|
||||||
version = "0.7.18"
|
version = "0.7.18"
|
||||||
@@ -1874,6 +1890,19 @@ dependencies = [
|
|||||||
"unicode-ident",
|
"unicode-ident",
|
||||||
]
|
]
|
||||||
|
|
||||||
|
[[package]]
|
||||||
|
name = "wasm-streams"
|
||||||
|
version = "0.5.0"
|
||||||
|
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
|
checksum = "9d1ec4f6517c9e11ae630e200b2b65d193279042e28edd4a2cda233e46670bbb"
|
||||||
|
dependencies = [
|
||||||
|
"futures-util",
|
||||||
|
"js-sys",
|
||||||
|
"wasm-bindgen",
|
||||||
|
"wasm-bindgen-futures",
|
||||||
|
"web-sys",
|
||||||
|
]
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
name = "web-sys"
|
name = "web-sys"
|
||||||
version = "0.3.94"
|
version = "0.3.94"
|
||||||
|
|||||||
+3
-1
@@ -13,7 +13,9 @@ tokio = { version = "1", features = ["full"] }
|
|||||||
serde = { version = "1", features = ["derive"] }
|
serde = { version = "1", features = ["derive"] }
|
||||||
serde_json = "1"
|
serde_json = "1"
|
||||||
jsonwebtoken = { version = "10.3.0", features = ["aws_lc_rs"] }
|
jsonwebtoken = { version = "10.3.0", features = ["aws_lc_rs"] }
|
||||||
reqwest = { version = "0.13.2", features = ["json"] }
|
reqwest = { version = "0.13.2", features = ["json", "stream"] }
|
||||||
once_cell = "1"
|
once_cell = "1"
|
||||||
dotenvy = "0.15"
|
dotenvy = "0.15"
|
||||||
thiserror = "2.0.18"
|
thiserror = "2.0.18"
|
||||||
|
tokio-stream = "0.1"
|
||||||
|
futures = "0.3"
|
||||||
+219
-67
@@ -1,6 +1,9 @@
|
|||||||
use crate::errors::OllamaError;
|
use crate::errors::OllamaError;
|
||||||
|
use axum::response::sse::Event;
|
||||||
|
use futures::StreamExt;
|
||||||
use reqwest::Client;
|
use reqwest::Client;
|
||||||
use serde_json::{Value, json};
|
use serde_json::{Value, json};
|
||||||
|
use tokio_stream::wrappers::ReceiverStream;
|
||||||
|
|
||||||
#[derive(Clone)]
|
#[derive(Clone)]
|
||||||
pub struct OllamaProvider {
|
pub struct OllamaProvider {
|
||||||
@@ -66,6 +69,99 @@ impl OllamaProvider {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
fn extract_completion_params<'a>(
|
||||||
|
&self,
|
||||||
|
body: &'a Value,
|
||||||
|
) -> Result<(&'a str, &'a str), OllamaError> {
|
||||||
|
let prompt = body
|
||||||
|
.get("prompt")
|
||||||
|
.and_then(|v| v.as_str())
|
||||||
|
.filter(|s| !s.trim().is_empty())
|
||||||
|
.ok_or(OllamaError::MissingPrompt)?;
|
||||||
|
|
||||||
|
let model = body
|
||||||
|
.get("model")
|
||||||
|
.and_then(|v| v.as_str())
|
||||||
|
.unwrap_or("llama3");
|
||||||
|
|
||||||
|
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 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 ─────────────────────────────────────────────────────
|
||||||
|
|
||||||
pub async fn list_models(&self) -> Result<Value, OllamaError> {
|
pub async fn list_models(&self) -> Result<Value, OllamaError> {
|
||||||
@@ -133,19 +229,10 @@ impl OllamaProvider {
|
|||||||
}
|
}
|
||||||
|
|
||||||
pub async fn completions(&self, body: Value) -> Result<Value, OllamaError> {
|
pub async fn completions(&self, body: Value) -> Result<Value, OllamaError> {
|
||||||
let prompt = body
|
let (prompt, model) = self.extract_completion_params(&body)?;
|
||||||
.get("prompt")
|
|
||||||
.and_then(|v| v.as_str())
|
|
||||||
.filter(|s| !s.trim().is_empty())
|
|
||||||
.ok_or(OllamaError::MissingPrompt)?;
|
|
||||||
|
|
||||||
let model = body
|
|
||||||
.get("model")
|
|
||||||
.and_then(|v| v.as_str())
|
|
||||||
.unwrap_or("llama3");
|
|
||||||
self.validate_model(model).await?;
|
self.validate_model(model).await?;
|
||||||
|
|
||||||
let ollama_payload = json!({
|
let payload = json!({
|
||||||
"model": model,
|
"model": model,
|
||||||
"prompt": prompt,
|
"prompt": prompt,
|
||||||
"stream": false,
|
"stream": false,
|
||||||
@@ -155,56 +242,80 @@ impl OllamaProvider {
|
|||||||
let res = self
|
let res = self
|
||||||
.client
|
.client
|
||||||
.post(format!("{}/api/generate", self.base_url))
|
.post(format!("{}/api/generate", self.base_url))
|
||||||
.json(&ollama_payload)
|
.json(&payload)
|
||||||
.send()
|
.send()
|
||||||
.await?
|
.await?
|
||||||
.json::<Value>()
|
.json::<Value>()
|
||||||
.await?;
|
.await?;
|
||||||
|
|
||||||
Ok(json!({
|
Ok(self.format_completion_response(&res))
|
||||||
|
}
|
||||||
|
|
||||||
|
pub async fn completions_stream(
|
||||||
|
&self,
|
||||||
|
body: Value,
|
||||||
|
) -> Result<ReceiverStream<Result<Event, OllamaError>>, OllamaError> {
|
||||||
|
let (prompt, model) = self.extract_completion_params(&body)?;
|
||||||
|
self.validate_model(model).await?;
|
||||||
|
|
||||||
|
let payload = json!({
|
||||||
|
"model": model,
|
||||||
|
"prompt": prompt,
|
||||||
|
"stream": true,
|
||||||
|
"options": Self::build_options(&body),
|
||||||
|
});
|
||||||
|
|
||||||
|
let mut byte_stream = self
|
||||||
|
.client
|
||||||
|
.post(format!("{}/api/generate", self.base_url))
|
||||||
|
.json(&payload)
|
||||||
|
.send()
|
||||||
|
.await?
|
||||||
|
.bytes_stream();
|
||||||
|
|
||||||
|
let (tx, rx) = tokio::sync::mpsc::channel(32);
|
||||||
|
|
||||||
|
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;
|
||||||
|
}
|
||||||
|
};
|
||||||
|
|
||||||
|
if let Ok(json) = serde_json::from_slice::<Value>(&chunk) {
|
||||||
|
let token = json.get("response").and_then(|v| v.as_str()).unwrap_or("");
|
||||||
|
let done = json.get("done").and_then(|v| v.as_bool()).unwrap_or(false);
|
||||||
|
|
||||||
|
// OpenAI-compatible SSE chunk
|
||||||
|
let event_data = serde_json::to_string(&json!({
|
||||||
"id": "cmpl-ollama",
|
"id": "cmpl-ollama",
|
||||||
"object": "text_completion",
|
"object": "text_completion",
|
||||||
"model": res.get("model"),
|
"choices": [{ "text": token, "index": 0, "finish_reason": null }],
|
||||||
"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,
|
|
||||||
}
|
|
||||||
}))
|
}))
|
||||||
|
.unwrap_or_default();
|
||||||
|
|
||||||
|
let _ = tx.send(Ok(Event::default().data(event_data))).await;
|
||||||
|
|
||||||
|
if done {
|
||||||
|
// Final [DONE] sentinel — matches OpenAI streaming protocol
|
||||||
|
let _ = tx.send(Ok(Event::default().data("[DONE]"))).await;
|
||||||
|
break;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
});
|
||||||
|
|
||||||
|
Ok(ReceiverStream::new(rx))
|
||||||
}
|
}
|
||||||
|
|
||||||
pub async fn chat_completions(&self, body: Value) -> Result<Value, OllamaError> {
|
pub async fn chat_completions(&self, body: Value) -> Result<Value, OllamaError> {
|
||||||
let model = body
|
let (model, messages) = self.extract_chat_params(&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);
|
|
||||||
}
|
|
||||||
|
|
||||||
self.validate_model(model).await?;
|
self.validate_model(model).await?;
|
||||||
|
|
||||||
let ollama_payload = json!({
|
let payload = json!({
|
||||||
"model": model,
|
"model": model,
|
||||||
"messages": messages,
|
"messages": messages,
|
||||||
"stream": false,
|
"stream": false,
|
||||||
@@ -214,35 +325,76 @@ impl OllamaProvider {
|
|||||||
let res = self
|
let res = self
|
||||||
.client
|
.client
|
||||||
.post(format!("{}/api/chat", self.base_url))
|
.post(format!("{}/api/chat", self.base_url))
|
||||||
.json(&ollama_payload)
|
.json(&payload)
|
||||||
.send()
|
.send()
|
||||||
.await?
|
.await?
|
||||||
.json::<Value>()
|
.json::<Value>()
|
||||||
.await?;
|
.await?;
|
||||||
|
|
||||||
Ok(json!({
|
Ok(self.format_chat_response(&res))
|
||||||
|
}
|
||||||
|
|
||||||
|
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 payload = json!({
|
||||||
|
"model": model,
|
||||||
|
"messages": messages,
|
||||||
|
"stream": true,
|
||||||
|
"options": Self::build_options(&body),
|
||||||
|
});
|
||||||
|
|
||||||
|
let mut byte_stream = self
|
||||||
|
.client
|
||||||
|
.post(format!("{}/api/chat", self.base_url))
|
||||||
|
.json(&payload)
|
||||||
|
.send()
|
||||||
|
.await?
|
||||||
|
.bytes_stream();
|
||||||
|
|
||||||
|
let (tx, rx) = tokio::sync::mpsc::channel(32);
|
||||||
|
|
||||||
|
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;
|
||||||
|
}
|
||||||
|
};
|
||||||
|
|
||||||
|
if let Ok(json) = serde_json::from_slice::<Value>(&chunk) {
|
||||||
|
let done = json.get("done").and_then(|v| v.as_bool()).unwrap_or(false);
|
||||||
|
|
||||||
|
let event_data = serde_json::to_string(&json!({
|
||||||
"id": "chatcmpl-ollama",
|
"id": "chatcmpl-ollama",
|
||||||
"object": "chat.completion",
|
"object": "chat.completion.chunk",
|
||||||
"model": res.get("model"),
|
|
||||||
"choices": [{
|
"choices": [{
|
||||||
"index": 0,
|
"index": 0,
|
||||||
"message": {
|
"delta": {
|
||||||
"role": res.get("message").and_then(|m| m.get("role")),
|
"role": json.get("message").and_then(|m| m.get("role")),
|
||||||
"content": res.get("message").and_then(|m| m.get("content")),
|
"content": json.get("message").and_then(|m| m.get("content")),
|
||||||
},
|
},
|
||||||
"finish_reason": res
|
"finish_reason": if done { json!("stop") } else { json!(null) },
|
||||||
.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),
|
|
||||||
}
|
|
||||||
}))
|
}))
|
||||||
|
.unwrap_or_default();
|
||||||
|
|
||||||
|
let _ = tx.send(Ok(Event::default().data(event_data))).await;
|
||||||
|
|
||||||
|
if done {
|
||||||
|
let _ = tx.send(Ok(Event::default().data("[DONE]"))).await;
|
||||||
|
break;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
});
|
||||||
|
|
||||||
|
Ok(ReceiverStream::new(rx))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
+71
-25
@@ -1,4 +1,11 @@
|
|||||||
use axum::{Json, extract::State};
|
use axum::{
|
||||||
|
Json,
|
||||||
|
extract::State,
|
||||||
|
response::{
|
||||||
|
IntoResponse, Response,
|
||||||
|
sse::{KeepAlive, Sse},
|
||||||
|
},
|
||||||
|
};
|
||||||
use serde_json::Value;
|
use serde_json::Value;
|
||||||
|
|
||||||
use crate::errors::OllamaError;
|
use crate::errors::OllamaError;
|
||||||
@@ -7,38 +14,77 @@ use crate::state::app_state::AppState;
|
|||||||
pub async fn completions(
|
pub async fn completions(
|
||||||
State(state): State<AppState>,
|
State(state): State<AppState>,
|
||||||
Json(body): Json<Value>,
|
Json(body): Json<Value>,
|
||||||
) -> Result<Json<Value>, (axum::http::StatusCode, String)> {
|
) -> Result<Response, (axum::http::StatusCode, String)> {
|
||||||
match state.ollama.completions(body).await {
|
let wants_stream = body
|
||||||
Ok(response) => Ok(Json(response)),
|
.get("stream")
|
||||||
Err(OllamaError::MissingPrompt) => Err((
|
.and_then(|v| v.as_bool())
|
||||||
axum::http::StatusCode::BAD_REQUEST,
|
.unwrap_or(false);
|
||||||
"prompt is required and cannot be empty".to_string(),
|
|
||||||
)),
|
if wants_stream {
|
||||||
Err(OllamaError::ModelNotFound(m)) => Err((
|
let stream = state
|
||||||
axum::http::StatusCode::UNPROCESSABLE_ENTITY,
|
.ollama
|
||||||
format!("model '{m}' is not available — run `ollama pull {m}` first"),
|
.completions_stream(body)
|
||||||
)),
|
.await
|
||||||
Err(e) => Err((axum::http::StatusCode::INTERNAL_SERVER_ERROR, e.to_string())),
|
.map_err(ollama_err)?;
|
||||||
|
|
||||||
|
Ok(Sse::new(stream)
|
||||||
|
.keep_alive(KeepAlive::default())
|
||||||
|
.into_response())
|
||||||
|
} else {
|
||||||
|
let response = state.ollama.completions(body).await.map_err(ollama_err)?;
|
||||||
|
|
||||||
|
Ok(Json(response).into_response())
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
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<Value>,
|
||||||
) -> Result<Json<Value>, (axum::http::StatusCode, String)> {
|
) -> Result<Response, (axum::http::StatusCode, String)> {
|
||||||
match state.ollama.chat_completions(body).await {
|
let wants_stream = body
|
||||||
Ok(response) => Ok(Json(response)),
|
.get("stream")
|
||||||
Err(OllamaError::MissingMessages) => Err((
|
.and_then(|v| v.as_bool())
|
||||||
|
.unwrap_or(false);
|
||||||
|
|
||||||
|
if wants_stream {
|
||||||
|
let stream = state
|
||||||
|
.ollama
|
||||||
|
.chat_completions_stream(body)
|
||||||
|
.await
|
||||||
|
.map_err(ollama_err)?;
|
||||||
|
|
||||||
|
Ok(Sse::new(stream)
|
||||||
|
.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) {
|
||||||
|
match e {
|
||||||
|
OllamaError::MissingPrompt => (
|
||||||
axum::http::StatusCode::BAD_REQUEST,
|
axum::http::StatusCode::BAD_REQUEST,
|
||||||
"messages must be a non-empty array with at least one user message".to_string(),
|
"prompt is required and cannot be empty".to_string(),
|
||||||
)),
|
),
|
||||||
Err(OllamaError::ModelNotFound(m)) => Err((
|
OllamaError::ModelNotFound(m) => (
|
||||||
axum::http::StatusCode::UNPROCESSABLE_ENTITY,
|
axum::http::StatusCode::UNPROCESSABLE_ENTITY,
|
||||||
format!("model '{m}' is not available — run `ollama pull {m}` first"),
|
format!("model '{m}' is not available — run `ollama pull {m}` first"),
|
||||||
)),
|
),
|
||||||
Err(OllamaError::Http(e)) => {
|
OllamaError::InvalidKeepAlive(v) => (
|
||||||
Err((axum::http::StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))
|
axum::http::StatusCode::BAD_REQUEST,
|
||||||
}
|
format!("invalid keep_alive '{v}'"),
|
||||||
Err(e) => Err((axum::http::StatusCode::BAD_REQUEST, e.to_string())),
|
),
|
||||||
|
OllamaError::MissingMessages => (
|
||||||
|
axum::http::StatusCode::BAD_REQUEST,
|
||||||
|
"messages array with at least one user message is required".to_string(),
|
||||||
|
),
|
||||||
|
OllamaError::Http(e) => (axum::http::StatusCode::INTERNAL_SERVER_ERROR, e.to_string()),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user