diff --git a/Cargo.lock b/Cargo.lock index cf3cd69..3f35dfd 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -11,6 +11,21 @@ dependencies = [ "memchr", ] +[[package]] +name = "android_system_properties" +version = "0.1.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "819e7219dbd41043ac279b19830f2efc897156490d7fd6ea916720117ee66311" +dependencies = [ + "libc", +] + +[[package]] +name = "anyhow" +version = "1.0.102" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7f202df86484c868dbad7eaa557ef785d5c66295e41b460ef922eca0723b842c" + [[package]] name = "assert-json-diff" version = "2.0.2" @@ -167,6 +182,7 @@ name = "chat" version = "0.1.0" dependencies = [ "axum", + "chrono", "dotenvy", "futures", "jsonwebtoken", @@ -178,9 +194,24 @@ dependencies = [ "tokio", "tokio-stream", "utoipa", + "uuid", "wiremock", ] +[[package]] +name = "chrono" +version = "0.4.44" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c673075a2e0e5f4a1dde27ce9dee1ea4558c7ffe648f576438a20ca1d2acc4b0" +dependencies = [ + "iana-time-zone", + "js-sys", + "num-traits", + "serde", + "wasm-bindgen", + "windows-link", +] + [[package]] name = "cmake" version = "0.1.58" @@ -313,6 +344,12 @@ version = "1.0.7" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "3f9eec918d3f24069decb9af1554cad7c880e2da24a9afd88aca000531ab82c1" +[[package]] +name = "foldhash" +version = "0.1.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d9c4f5dac5e15c24eb999c26181a6ca40b39fe946cbe4c263c7209467bc83af2" + [[package]] name = "form_urlencoded" version = "1.2.2" @@ -438,11 +475,24 @@ dependencies = [ "cfg-if", "js-sys", "libc", - "r-efi", + "r-efi 5.3.0", "wasip2", "wasm-bindgen", ] +[[package]] +name = "getrandom" +version = "0.4.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0de51e6874e94e7bf76d726fc5d13ba782deca734ff60d5bb2fb2607c7406555" +dependencies = [ + "cfg-if", + "libc", + "r-efi 6.0.0", + "wasip2", + "wasip3", +] + [[package]] name = "h2" version = "0.4.13" @@ -462,12 +512,27 @@ dependencies = [ "tracing", ] +[[package]] +name = "hashbrown" +version = "0.15.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9229cfe53dfd69f0609a49f65461bd93001ea1ef889cd5529dd176593f5338a1" +dependencies = [ + "foldhash", +] + [[package]] name = "hashbrown" version = "0.16.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "841d1cc9bed7f9236f321df977030373f4a4163ae1a7dbfe1a51a2c1a51d9100" +[[package]] +name = "heck" +version = "0.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2304e00983f87ffb38b55b444b5e3b60a884b5d30c0fca7d82fe33449bbe55ea" + [[package]] name = "hermit-abi" version = "0.5.2" @@ -582,6 +647,30 @@ dependencies = [ "windows-registry", ] +[[package]] +name = "iana-time-zone" +version = "0.1.65" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e31bc9ad994ba00e440a8aa5c9ef0ec67d5cb5e5cb0cc7f8b744a35b389cc470" +dependencies = [ + "android_system_properties", + "core-foundation-sys", + "iana-time-zone-haiku", + "js-sys", + "log", + "wasm-bindgen", + "windows-core", +] + +[[package]] +name = "iana-time-zone-haiku" +version = "0.1.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f31827a206f56af32e590ba56d5d2d085f558508192593743f16b2306495269f" +dependencies = [ + "cc", +] + [[package]] name = "icu_collections" version = "2.2.0" @@ -664,6 +753,12 @@ dependencies = [ "zerovec", ] +[[package]] +name = "id-arena" +version = "2.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3d3067d79b975e8844ca9eb072e16b31c3c1c36928edf9c6789548c524d0d954" + [[package]] name = "idna" version = "1.1.0" @@ -692,7 +787,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "45a8a2b9cb3e0b0c1803dbb0758ffac5de2f425b23c28f518faabd9d805342ff" dependencies = [ "equivalent", - "hashbrown", + "hashbrown 0.16.1", "serde", "serde_core", ] @@ -808,6 +903,12 @@ version = "1.5.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "bbd2bcb4c963f2ddae06a2efc7e9f3591312473c50c6685e1f298068316e66fe" +[[package]] +name = "leb128fmt" +version = "0.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "09edd9e8b54e49e587e4f6295a7d29c3ea94d469cb40ab8ca70b288248a81db2" + [[package]] name = "libc" version = "0.2.184" @@ -995,6 +1096,16 @@ dependencies = [ "zerocopy", ] +[[package]] +name = "prettyplease" +version = "0.2.37" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "479ca8adacdd7ce8f1fb39ce9ecccbfe93a3f1344b3d0d97f20bc0196208f62b" +dependencies = [ + "proc-macro2", + "syn", +] + [[package]] name = "proc-macro2" version = "1.0.106" @@ -1075,6 +1186,12 @@ version = "5.3.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "69cdb34c158ceb288df11e18b4bd39de994f6657d83847bdffdbd7f346754b0f" +[[package]] +name = "r-efi" +version = "6.0.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f8dcc9c7d52a811697d2151c701e0d08956f92b0e24136cf4cf27b57a6a0d9bf" + [[package]] name = "rand" version = "0.9.2" @@ -1348,6 +1465,12 @@ dependencies = [ "libc", ] +[[package]] +name = "semver" +version = "1.0.28" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8a7852d02fc848982e0c167ef163aaff9cd91dc640ba85e263cb1ce46fae51cd" + [[package]] name = "serde" version = "1.0.228" @@ -1774,6 +1897,12 @@ version = "1.0.24" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "e6e4313cd5fcd3dad5cafa179702e2b244f760991f45397d14d4ebf38247da75" +[[package]] +name = "unicode-xid" +version = "0.2.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ebc1c04c71510c7f702b52b7c350734c9ff1295c464a03335b00bb84fc54f853" + [[package]] name = "untrusted" version = "0.7.1" @@ -1828,6 +1957,18 @@ dependencies = [ "syn", ] +[[package]] +name = "uuid" +version = "1.23.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5ac8b6f42ead25368cf5b098aeb3dc8a1a2c05a3eee8a9a1a68c640edbfc79d9" +dependencies = [ + "getrandom 0.4.2", + "js-sys", + "serde_core", + "wasm-bindgen", +] + [[package]] name = "walkdir" version = "2.5.0" @@ -1862,6 +2003,15 @@ dependencies = [ "wit-bindgen", ] +[[package]] +name = "wasip3" +version = "0.4.0+wasi-0.3.0-rc-2026-01-06" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5428f8bf88ea5ddc08faddef2ac4a67e390b88186c703ce6dbd955e1c145aca5" +dependencies = [ + "wit-bindgen", +] + [[package]] name = "wasm-bindgen" version = "0.2.117" @@ -1917,6 +2067,28 @@ dependencies = [ "unicode-ident", ] +[[package]] +name = "wasm-encoder" +version = "0.244.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "990065f2fe63003fe337b932cfb5e3b80e0b4d0f5ff650e6985b1048f62c8319" +dependencies = [ + "leb128fmt", + "wasmparser", +] + +[[package]] +name = "wasm-metadata" +version = "0.244.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "bb0e353e6a2fbdc176932bbaab493762eb1255a7900fe0fea1a2f96c296cc909" +dependencies = [ + "anyhow", + "indexmap", + "wasm-encoder", + "wasmparser", +] + [[package]] name = "wasm-streams" version = "0.5.0" @@ -1930,6 +2102,18 @@ dependencies = [ "web-sys", ] +[[package]] +name = "wasmparser" +version = "0.244.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "47b807c72e1bac69382b3a6fb3dbe8ea4c0ed87ff5629b8685ae6b9a611028fe" +dependencies = [ + "bitflags", + "hashbrown 0.15.5", + "indexmap", + "semver", +] + [[package]] name = "web-sys" version = "0.3.94" @@ -1968,6 +2152,41 @@ dependencies = [ "windows-sys 0.61.2", ] +[[package]] +name = "windows-core" +version = "0.62.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b8e83a14d34d0623b51dce9581199302a221863196a1dde71a7663a4c2be9deb" +dependencies = [ + "windows-implement", + "windows-interface", + "windows-link", + "windows-result", + "windows-strings", +] + +[[package]] +name = "windows-implement" +version = "0.60.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "053e2e040ab57b9dc951b72c264860db7eb3b0200ba345b4e4c3b14f67855ddf" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "windows-interface" +version = "0.59.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3f316c4a2570ba26bbec722032c4099d8c8bc095efccdc15688708623367e358" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + [[package]] name = "windows-link" version = "0.2.1" @@ -2253,6 +2472,88 @@ name = "wit-bindgen" version = "0.51.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "d7249219f66ced02969388cf2bb044a09756a083d0fab1e566056b04d9fbcaa5" +dependencies = [ + "wit-bindgen-rust-macro", +] + +[[package]] +name = "wit-bindgen-core" +version = "0.51.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ea61de684c3ea68cb082b7a88508a8b27fcc8b797d738bfc99a82facf1d752dc" +dependencies = [ + "anyhow", + "heck", + "wit-parser", +] + +[[package]] +name = "wit-bindgen-rust" +version = "0.51.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b7c566e0f4b284dd6561c786d9cb0142da491f46a9fbed79ea69cdad5db17f21" +dependencies = [ + "anyhow", + "heck", + "indexmap", + "prettyplease", + "syn", + "wasm-metadata", + "wit-bindgen-core", + "wit-component", +] + +[[package]] +name = "wit-bindgen-rust-macro" +version = "0.51.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0c0f9bfd77e6a48eccf51359e3ae77140a7f50b1e2ebfe62422d8afdaffab17a" +dependencies = [ + "anyhow", + "prettyplease", + "proc-macro2", + "quote", + "syn", + "wit-bindgen-core", + "wit-bindgen-rust", +] + +[[package]] +name = "wit-component" +version = "0.244.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9d66ea20e9553b30172b5e831994e35fbde2d165325bec84fc43dbf6f4eb9cb2" +dependencies = [ + "anyhow", + "bitflags", + "indexmap", + "log", + "serde", + "serde_derive", + "serde_json", + "wasm-encoder", + "wasm-metadata", + "wasmparser", + "wit-parser", +] + +[[package]] +name = "wit-parser" +version = "0.244.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ecc8ac4bc1dc3381b7f59c34f00b67e18f910c2c0f50015669dde7def656a736" +dependencies = [ + "anyhow", + "id-arena", + "indexmap", + "log", + "semver", + "serde", + "serde_derive", + "serde_json", + "unicode-xid", + "wasmparser", +] [[package]] name = "writeable" diff --git a/Cargo.toml b/Cargo.toml index 02e9182..8fb2d63 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -19,4 +19,6 @@ once_cell = "1" dotenvy = "0.15" thiserror = "2.0.18" tokio-stream = "0.1" -futures = "0.3" \ No newline at end of file +futures = "0.3" +chrono = { version = "0.4.44", features = ["serde"] } +uuid = { version = "1", features = ["v4", "serde"] } \ No newline at end of file diff --git a/src/dto/api.rs b/src/dto/api.rs index e3c0fea..fdd8ca4 100644 --- a/src/dto/api.rs +++ b/src/dto/api.rs @@ -3,21 +3,6 @@ use utoipa::ToSchema; use crate::errors::OllamaError; -#[derive(Serialize, Deserialize, ToSchema)] -pub struct ChatRequest { - pub model: String, - pub prompt: Option, - pub messages: Option>, - #[serde(default)] - pub stream: bool, - pub temperature: Option, - pub top_p: Option, - pub max_tokens: Option, - pub stop: Option>, - pub system: Option, - pub keep_alive: Option, -} - #[derive(Debug, Serialize, Deserialize, ToSchema)] #[serde(rename_all = "lowercase")] pub enum Role { @@ -37,6 +22,8 @@ pub struct Message { pub content: String, } +// --------------------------- + #[derive(Debug, Serialize, Deserialize, ToSchema)] pub struct ModelsResponse { pub models: Vec, @@ -91,3 +78,80 @@ pub struct UnloadModelResponse { pub model: String, pub status: String, } + +#[derive(Debug, Deserialize, Serialize, ToSchema)] +pub struct BaseLLMRequest { + pub model: String, + + #[serde(default)] + pub stream: bool, + + pub temperature: Option, + pub top_p: Option, + + // Ollama-native + pub top_k: Option, + pub repeat_penalty: Option, + pub seed: Option, + + pub num_ctx: Option, + pub num_predict: Option, + + pub stop: Option>, + + pub keep_alive: Option, +} + +#[derive(Debug, Deserialize, Serialize, ToSchema)] +pub struct CompletionRequest { + #[serde(flatten)] + pub base: BaseLLMRequest, + + pub prompt: String, +} + +#[derive(Debug, Serialize, Deserialize, ToSchema)] +#[serde(rename_all = "snake_case")] +pub enum CompletionObject { + TextCompletion, +} + +#[derive(Debug, Serialize, Deserialize, ToSchema)] +#[serde(rename_all = "snake_case")] +pub enum FinishReason { + Stop, + Length, + ContentFilter, + ToolCalls, +} + +#[derive(Debug, Serialize, Deserialize, ToSchema)] +pub struct CompletionResponse { + pub id: String, + pub object: CompletionObject, + pub created: u64, + pub model: String, + pub choices: Vec, + pub usage: Usage, +} + +#[derive(Debug, Serialize, Deserialize, ToSchema)] +pub struct Choice { + pub text: String, + pub index: u32, + pub finish_reason: FinishReason, +} + +#[derive(Debug, Serialize, Deserialize, ToSchema)] +pub struct Usage { + pub prompt_tokens: u32, + pub completion_tokens: u32, + pub total_tokens: u32, +} + +#[derive(Debug, Serialize, Deserialize, ToSchema)] +pub struct CompletionChunk { + pub id: String, + pub object: String, + pub choices: Vec, +} diff --git a/src/dto/ollama.rs b/src/dto/ollama.rs index 06b69a5..ce870b5 100644 --- a/src/dto/ollama.rs +++ b/src/dto/ollama.rs @@ -1,8 +1,6 @@ use serde::{Deserialize, Serialize}; use utoipa::ToSchema; -use crate::errors::OllamaError; - #[derive(Debug, Serialize, Deserialize)] pub struct OllamaModels { pub models: Vec, @@ -25,3 +23,39 @@ pub struct OllamaModelDetails { pub parameter_size: Option, pub quantization_level: Option, } + +#[derive(Debug, Serialize, Deserialize)] +pub struct OllamaOptions { + pub temperature: Option, + pub top_p: Option, + pub top_k: Option, + pub repeat_penalty: Option, + pub seed: Option, + + pub num_ctx: Option, + pub num_predict: Option, +} + +#[derive(Debug, Serialize)] +pub struct OllamaGenerateRequest<'a> { + pub model: &'a str, + pub prompt: &'a str, + pub stream: bool, + pub options: OllamaOptions, +} + +#[derive(Debug, Serialize, Deserialize)] +pub struct OllamaGenerateResponse { + pub model: String, + pub created_at: Option, + pub response: String, + pub done: bool, + + #[serde(default)] + pub context: Option>, + + pub total_duration: Option, + pub load_duration: Option, + pub prompt_eval_count: Option, + pub eval_count: Option, +} diff --git a/src/providers/ollama/client.rs b/src/providers/ollama/client.rs index 825611e..6e0684c 100644 --- a/src/providers/ollama/client.rs +++ b/src/providers/ollama/client.rs @@ -1,5 +1,6 @@ use crate::dto::{api, ollama}; use crate::errors::OllamaError; +use axum::Json; use axum::response::sse::Event; use futures::StreamExt; use reqwest::Client; @@ -36,48 +37,16 @@ impl OllamaProvider { Ok(res.models.iter().any(|m| m.name == model)) } - // fn build_options(body: &Value) -> Value { - // json!({ - // "temperature": body.get("temperature"), - // "top_p": body.get("top_p"), - // "num_predict": body.get("max_tokens"), - // }) - // } + fn extract_completion_params<'a>( + &self, + body: &'a api::CompletionRequest, + ) -> Result<(&'a str, &'a str), OllamaError> { + let prompt = body.prompt.trim(); - // async fn validate_model(&self, model: &str) -> Result<(), OllamaError> { - // let available = self.list_models().await?; - // let exists = available - // .get("models") - // .and_then(|m| m.as_array()) - // .map(|arr| { - // arr.iter() - // .any(|m| m.get("name").and_then(|n| n.as_str()) == Some(model)) - // }) - // .unwrap_or(false); + let model = body.base.model.as_str(); - // if !exists { - // return Err(OllamaError::ModelNotFound(model.to_string())); - // } - // Ok(()) - // } - - // 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()) - // .ok_or(OllamaError::MissingModel)?; - - // Ok((prompt, model)) - // } + Ok((prompt, model)) + } // fn format_completion_response(&self, res: &Value) -> Value { // json!({ @@ -209,7 +178,7 @@ impl OllamaProvider { pub async fn unload_model(&self, model: &str) -> Result { let url = format!("{}/api/generate", self.base_url); - + let exists = self.model_exists(model).await?; if !exists { return Err(OllamaError::ModelNotFound(model.to_string())); @@ -237,88 +206,131 @@ impl OllamaProvider { }) } - // pub async fn completions(&self, body: Value) -> Result { - // let (prompt, model) = self.extract_completion_params(&body)?; - // self.validate_model(model).await?; + pub async fn completions( + &self, + body: &api::CompletionRequest, + ) -> Result { + let url = format!("{}/api/generate", self.base_url); - // let payload = json!({ - // "model": model, - // "prompt": prompt, - // "stream": false, - // "options": Self::build_options(&body), - // }); + let (prompt, model) = self.extract_completion_params(body)?; - // let res = self - // .client - // .post(format!("{}/api/generate", self.base_url)) - // .json(&payload) - // .send() - // .await? - // .json::() - // .await?; + if prompt.is_empty() { + return Err(OllamaError::MissingPrompt); + } - // Ok(self.format_completion_response(&res)) - // } + if model.is_empty() { + return Err(OllamaError::MissingModel); + } - // pub async fn completions_stream( - // &self, - // body: Value, - // ) -> Result>, OllamaError> { - // let (prompt, model) = self.extract_completion_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, - // "prompt": prompt, - // "stream": true, - // "options": Self::build_options(&body), - // }); + let options = ollama::OllamaOptions::from(body); - // let mut byte_stream = self - // .client - // .post(format!("{}/api/generate", self.base_url)) - // .json(&payload) - // .send() - // .await? - // .bytes_stream(); + let payload = ollama::OllamaGenerateRequest { + model, + prompt, + stream: false, + options, + }; - // let (tx, rx) = tokio::sync::mpsc::channel(32); + let res = self + .client + .post(url) + .json(&payload) + .send() + .await? + .json::() + .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::CompletionResponse::from(res)) + } - // if let Ok(json) = serde_json::from_slice::(&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); + pub async fn completions_stream( + &self, + body: &api::CompletionRequest, + ) -> Result>, OllamaError> { + let url = format!("{}/api/generate", self.base_url); - // // OpenAI-compatible SSE chunk - // let event_data = serde_json::to_string(&json!({ - // "id": "cmpl-ollama", - // "object": "text_completion", - // "choices": [{ "text": token, "index": 0, "finish_reason": null }], - // })) - // .unwrap_or_default(); + let (prompt, model) = self.extract_completion_params(body)?; - // let _ = tx.send(Ok(Event::default().data(event_data))).await; + if prompt.is_empty() { + return Err(OllamaError::MissingPrompt); + } - // if done { - // // Final [DONE] sentinel — matches OpenAI streaming protocol - // 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); + + let payload = ollama::OllamaGenerateRequest { + model, + prompt, + 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 { + 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; + } + }; + + // 🔥 IMPORTANT: typed deserialization + let parsed: ollama::OllamaGenerateResponse = match serde_json::from_slice(&chunk) { + Ok(v) => v, + Err(_) => continue, + }; + + // map → OpenAI chunk + let event_data = serde_json::to_string(&api::CompletionChunk { + id: "cmpl-ollama".to_string(), + object: "text_completion".to_string(), + choices: vec![api::Choice { + text: parsed.response, + index: 0, + finish_reason: if parsed.done { + api::FinishReason::Stop + } else { + api::FinishReason::Length + }, + }], + }) + .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)) + } // pub async fn chat_completions(&self, body: Value) -> Result { // let (model, messages) = self.extract_chat_params(&body)?; diff --git a/src/providers/ollama/mapper.rs b/src/providers/ollama/mapper.rs index b7e1052..87ccb5a 100644 --- a/src/providers/ollama/mapper.rs +++ b/src/providers/ollama/mapper.rs @@ -1,5 +1,8 @@ use crate::dto::{api, ollama}; +use chrono::Utc; +use uuid::Uuid; + impl From for api::ModelInfo { fn from(m: ollama::OllamaModel) -> Self { Self { @@ -14,3 +17,40 @@ impl From 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 for api::CompletionResponse { + fn from(res: ollama::OllamaGenerateResponse) -> Self { + Self { + id: Uuid::new_v4().to_string(), + object: api::CompletionObject::TextCompletion, + model: res.model, + created: Utc::now().timestamp() as u64, + + choices: vec![api::Choice { + text: res.response, + index: 0, + finish_reason: api::FinishReason::Stop, + }], + + usage: api::Usage { + prompt_tokens: res.prompt_eval_count.unwrap_or(0), + completion_tokens: res.eval_count.unwrap_or(0), + total_tokens: res.prompt_eval_count.unwrap_or(0) + res.eval_count.unwrap_or(0), + }, + } + } +} diff --git a/src/routes/v1/chat.rs b/src/routes/v1/chat.rs index 2fa030f..646546a 100644 --- a/src/routes/v1/chat.rs +++ b/src/routes/v1/chat.rs @@ -8,6 +8,7 @@ use axum::{ }; use serde_json::Value; +use crate::dto::api; use crate::errors::OllamaError; use crate::state::app_state::AppState; @@ -18,31 +19,28 @@ use crate::state::app_state::AppState; (status = 200, description = "Chat completion", body = Value), ) )] -// pub async fn completions( -// State(state): State, -// Json(body): Json, -// ) -> Result { -// let wants_stream = body -// .get("stream") -// .and_then(|v| v.as_bool()) -// .unwrap_or(false); +pub async fn completions( + State(state): State, + Json(body): Json, +) -> Result { + let wants_stream = body.base.stream; -// if wants_stream { -// let stream = state -// .ollama -// .completions_stream(body) -// .await -// .map_err(ollama_err)?; + if wants_stream { + let stream = state + .ollama + .completions_stream(&body) + .await + .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(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()) -// } -// } + Ok(Json(response).into_response()) + } +} // pub async fn chat_completions( // State(state): State, diff --git a/src/routes/v1/mod.rs b/src/routes/v1/mod.rs index 8c7246b..a1aaf47 100644 --- a/src/routes/v1/mod.rs +++ b/src/routes/v1/mod.rs @@ -11,11 +11,12 @@ use axum::{Router, middleware, routing::get, routing::post}; // } pub fn protected_router() -> Router { - Router::new().route("/models", get(models::list_models)) - // .route("/completions", post(chat::completions)) - // .route("/chat/completions", post(chat::chat_completions)) - .route("/models/{model}/load", post(models::load_model)) - .route("/models/{model}/unload", post(models::unload_model)) + Router::new() + .route("/models", get(models::list_models)) + .route("/completions", post(chat::completions)) + // .route("/chat/completions", post(chat::chat_completions)) + .route("/models/{model}/load", post(models::load_model)) + .route("/models/{model}/unload", post(models::unload_model)) } pub fn router() -> Router { diff --git a/src/routes/v1/models.rs b/src/routes/v1/models.rs index d8a2cfb..32db1a9 100644 --- a/src/routes/v1/models.rs +++ b/src/routes/v1/models.rs @@ -37,7 +37,6 @@ pub async fn unload_model( State(state): State, Path(model): Path, ) -> Result, (axum::http::StatusCode, String)> { - let response = state .ollama .unload_model(&model)