reafctor: all code without stream
This commit is contained in:
@@ -0,0 +1,23 @@
|
||||
use serde::Deserialize;
|
||||
use std::collections::HashMap;
|
||||
|
||||
#[derive(Debug, Clone, Default, Deserialize)]
|
||||
pub struct KeycloakClaims {
|
||||
pub sub: String,
|
||||
pub preferred_username: Option<String>,
|
||||
pub exp: usize,
|
||||
pub iss: String,
|
||||
pub _aud: Option<Vec<String>>,
|
||||
pub realm_access: Option<RealmAccess>,
|
||||
pub resource_access: HashMap<String, ResourceAccess>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Deserialize)]
|
||||
pub struct RealmAccess {
|
||||
pub roles: Vec<String>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Deserialize)]
|
||||
pub struct ResourceAccess {
|
||||
pub roles: Vec<String>,
|
||||
}
|
||||
@@ -0,0 +1,40 @@
|
||||
use thiserror::Error;
|
||||
|
||||
#[derive(Debug, Error)]
|
||||
pub enum AuthError {
|
||||
#[error("invalid authorization header")]
|
||||
InvalidHeader,
|
||||
|
||||
#[error("missing 'kid' in token header")]
|
||||
MissingKid,
|
||||
|
||||
#[error("invalid JWKS structure")]
|
||||
InvalidJwks,
|
||||
|
||||
#[error("no matching key found for kid")]
|
||||
JwkNotFound,
|
||||
|
||||
#[error("missing RSA modulus")]
|
||||
MissingModulus,
|
||||
|
||||
#[error("missing RSA exponent")]
|
||||
MissingExponent,
|
||||
|
||||
#[error("invalid decoding key")]
|
||||
InvalidDecodingKey,
|
||||
|
||||
#[error("token validation failed")]
|
||||
TokenValidationFailed,
|
||||
|
||||
#[error("invalid or expired token")]
|
||||
InvalidToken,
|
||||
|
||||
#[error("failed to fetch JWKS")]
|
||||
JwksFetchFailed,
|
||||
|
||||
#[error("failed to refresh JWKS")]
|
||||
JwksRefreshFailed,
|
||||
|
||||
#[error(transparent)]
|
||||
Reqwest(#[from] reqwest::Error),
|
||||
}
|
||||
@@ -0,0 +1,59 @@
|
||||
use super::errors::AuthError;
|
||||
|
||||
use once_cell::sync::Lazy;
|
||||
use serde_json::Value;
|
||||
use std::env;
|
||||
use std::sync::Arc;
|
||||
use std::time::{Duration, Instant};
|
||||
use tokio::sync::RwLock;
|
||||
|
||||
#[derive(Clone)]
|
||||
struct JwksCache {
|
||||
jwks: Value,
|
||||
last_fetched: Instant,
|
||||
}
|
||||
|
||||
static JWK_CACHE: Lazy<Arc<RwLock<Option<JwksCache>>>> = Lazy::new(|| Arc::new(RwLock::new(None)));
|
||||
|
||||
static JWKS_URL: Lazy<String> = Lazy::new(|| env::var("JWKS_URL").expect("JWKS_URL not set"));
|
||||
|
||||
async fn fetch_jwks() -> Result<Value, AuthError> {
|
||||
let jwks = reqwest::get(JWKS_URL.as_str())
|
||||
.await?
|
||||
.json::<Value>()
|
||||
.await?;
|
||||
|
||||
Ok(jwks)
|
||||
}
|
||||
|
||||
pub async fn refresh_jwks() -> Result<Value, AuthError> {
|
||||
let jwks = fetch_jwks().await?;
|
||||
|
||||
let mut write = JWK_CACHE.write().await;
|
||||
|
||||
*write = Some(JwksCache {
|
||||
jwks: jwks.clone(),
|
||||
last_fetched: Instant::now(),
|
||||
});
|
||||
|
||||
Ok(jwks)
|
||||
}
|
||||
|
||||
pub async fn get_jwks() -> Result<Value, AuthError> {
|
||||
let ttl = Duration::from_secs(3600); // 1 hour
|
||||
|
||||
{
|
||||
// Read lock first (fast path)
|
||||
let read = JWK_CACHE.read().await;
|
||||
|
||||
if let Some(cache) = read
|
||||
.as_ref()
|
||||
.filter(|cache| cache.last_fetched.elapsed() < ttl)
|
||||
{
|
||||
return Ok(cache.jwks.clone());
|
||||
}
|
||||
}
|
||||
|
||||
// Expired or empty → refresh
|
||||
refresh_jwks().await
|
||||
}
|
||||
@@ -0,0 +1,4 @@
|
||||
pub mod claims;
|
||||
pub mod errors;
|
||||
pub mod jwks;
|
||||
pub mod validator;
|
||||
@@ -0,0 +1,64 @@
|
||||
use super::claims::KeycloakClaims;
|
||||
use super::errors::AuthError;
|
||||
|
||||
use jsonwebtoken::{Algorithm, DecodingKey, Validation, decode, decode_header};
|
||||
use once_cell::sync::Lazy;
|
||||
use serde_json::Value;
|
||||
use std::env;
|
||||
|
||||
static ISSUER: Lazy<String> = Lazy::new(|| env::var("ISSUER").expect("ISSUER not set"));
|
||||
|
||||
fn validate_token(token: &str, jwks: &Value) -> Result<KeycloakClaims, AuthError> {
|
||||
// 1. Decode header
|
||||
let header = decode_header(token).map_err(|_| AuthError::InvalidHeader)?;
|
||||
|
||||
let kid = header.kid.ok_or(AuthError::MissingKid)?;
|
||||
|
||||
// 2. Find matching key
|
||||
let keys = jwks["keys"].as_array().ok_or(AuthError::InvalidJwks)?;
|
||||
|
||||
let key = keys
|
||||
.iter()
|
||||
.find(|k| k["kid"] == kid)
|
||||
.ok_or(AuthError::InvalidDecodingKey)?;
|
||||
|
||||
// 3. Extract RSA components
|
||||
let n = key["n"].as_str().ok_or(AuthError::MissingModulus)?;
|
||||
let e = key["e"].as_str().ok_or(AuthError::MissingExponent)?;
|
||||
|
||||
let decoding_key =
|
||||
DecodingKey::from_rsa_components(n, e).map_err(|_| AuthError::JwkNotFound)?;
|
||||
|
||||
// 4. Setup validation rules
|
||||
let mut validation = Validation::new(Algorithm::RS256);
|
||||
|
||||
validation.set_issuer(&[ISSUER.as_str()]);
|
||||
|
||||
validation.validate_exp = true;
|
||||
validation.validate_aud = false;
|
||||
|
||||
// 5. Decode & verify
|
||||
let token_data = decode::<KeycloakClaims>(token, &decoding_key, &validation)
|
||||
.map_err(|_| AuthError::TokenValidationFailed)?;
|
||||
|
||||
Ok(token_data.claims)
|
||||
}
|
||||
|
||||
pub async fn authenticate_jwt(token: &str) -> Result<KeycloakClaims, AuthError> {
|
||||
let jwks = super::jwks::get_jwks()
|
||||
.await
|
||||
.map_err(|_| AuthError::JwksFetchFailed)?;
|
||||
|
||||
match validate_token(token, &jwks) {
|
||||
Ok(claims) => Ok(claims),
|
||||
|
||||
Err(_) => {
|
||||
// one retry with refresh
|
||||
let fresh = super::jwks::refresh_jwks()
|
||||
.await
|
||||
.map_err(|_| AuthError::JwksRefreshFailed)?;
|
||||
|
||||
validate_token(token, &fresh).map_err(|_| AuthError::InvalidToken)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -1 +1,2 @@
|
||||
pub mod keycloak;
|
||||
pub mod ollama;
|
||||
|
||||
+149
-265
@@ -1,10 +1,8 @@
|
||||
use super::errors::OllamaError;
|
||||
use crate::dto::{api, ollama};
|
||||
use axum::response::sse::Event;
|
||||
use crate::providers::ollama;
|
||||
use crate::providers::ollama::errors::LlmError;
|
||||
|
||||
use futures::StreamExt;
|
||||
use reqwest::Client;
|
||||
use serde_json::json;
|
||||
use tokio_stream::wrappers::ReceiverStream;
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct OllamaProvider {
|
||||
@@ -22,7 +20,7 @@ impl OllamaProvider {
|
||||
|
||||
// ── private helpers ──────────────────────────────────────────────────────
|
||||
|
||||
async fn model_exists(&self, model: &str) -> Result<bool, OllamaError> {
|
||||
async fn model_exists(&self, model: &str) -> Result<bool, LlmError> {
|
||||
let url = format!("{}/api/tags", self.base_url);
|
||||
|
||||
let res = self
|
||||
@@ -30,61 +28,43 @@ impl OllamaProvider {
|
||||
.get(url)
|
||||
.send()
|
||||
.await?
|
||||
.json::<ollama::OllamaModels>()
|
||||
.json::<ollama::types::OllamaModels>()
|
||||
.await?;
|
||||
|
||||
Ok(res.models.iter().any(|m| m.name == model))
|
||||
}
|
||||
|
||||
fn has_user_message(&self, messages: &[api::Message]) -> bool {
|
||||
messages.iter().any(|m| matches!(m.role, api::Role::User))
|
||||
fn has_user_message(&self, messages: &[ollama::types::OllamaMessage]) -> bool {
|
||||
messages
|
||||
.iter()
|
||||
.any(|m| matches!(m.role, ollama::types::OllamaRole::User))
|
||||
}
|
||||
|
||||
fn extract_completion_params<'a>(
|
||||
&self,
|
||||
body: &'a api::CompletionRequest,
|
||||
) -> Result<(&'a str, &'a str), OllamaError> {
|
||||
let prompt = body.prompt.trim();
|
||||
// pub fn validate_keep_alive(&self, s: &str) -> Result<(), LlmError> {
|
||||
// let s = s.trim();
|
||||
|
||||
let model = body.base.model.as_str();
|
||||
// if s == "-1" || s.parse::<u64>().is_ok() {
|
||||
// return Ok(());
|
||||
// }
|
||||
|
||||
Ok((prompt, model))
|
||||
}
|
||||
// let split = s
|
||||
// .find(|c: char| c.is_alphabetic())
|
||||
// .ok_or_else(|| LlmError::InvalidKeepAlive(s.to_string()))?;
|
||||
|
||||
fn extract_chat_params<'a>(
|
||||
&self,
|
||||
body: &'a api::ChatRequest,
|
||||
) -> Result<(&'a [api::Message], &'a str), OllamaError> {
|
||||
let model = body.base.model.as_str();
|
||||
// let (num, unit) = s.split_at(split);
|
||||
|
||||
Ok((&body.messages, model))
|
||||
}
|
||||
// num.parse::<u64>()
|
||||
// .map_err(|_| LlmError::InvalidKeepAlive(s.to_string()))?;
|
||||
|
||||
pub fn parse_keep_alive(&self, s: &str) -> Result<(), OllamaError> {
|
||||
let s = s.trim();
|
||||
|
||||
if s == "-1" || s.parse::<u64>().is_ok() {
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let split = s
|
||||
.find(|c: char| c.is_alphabetic())
|
||||
.ok_or_else(|| OllamaError::InvalidKeepAlive(s.to_string()))?;
|
||||
|
||||
let (num, unit) = s.split_at(split);
|
||||
|
||||
num.parse::<u64>()
|
||||
.map_err(|_| OllamaError::InvalidKeepAlive(s.to_string()))?;
|
||||
|
||||
match unit {
|
||||
"s" | "m" | "h" => Ok(()),
|
||||
_ => Err(OllamaError::InvalidKeepAlive(s.to_string())),
|
||||
}
|
||||
}
|
||||
// match unit {
|
||||
// "s" | "m" | "h" => Ok(()),
|
||||
// _ => Err(LlmError::InvalidKeepAlive(s.to_string())),
|
||||
// }
|
||||
// }
|
||||
|
||||
// // ── public endpoints ─────────────────────────────────────────────────────
|
||||
|
||||
pub async fn list_models(&self) -> Result<api::ModelsResponse, OllamaError> {
|
||||
pub async fn list_models(&self) -> Result<super::types::OllamaModels, LlmError> {
|
||||
let url = format!("{}/api/tags", self.base_url);
|
||||
|
||||
let res = self
|
||||
@@ -92,156 +72,83 @@ impl OllamaProvider {
|
||||
.get(url)
|
||||
.send()
|
||||
.await?
|
||||
.json::<ollama::OllamaModels>()
|
||||
.json::<ollama::types::OllamaModels>()
|
||||
.await?;
|
||||
|
||||
let models = res.models.into_iter().map(api::ModelInfo::from).collect();
|
||||
|
||||
Ok(api::ModelsResponse { models })
|
||||
}
|
||||
|
||||
pub async fn load_model(
|
||||
&self,
|
||||
model: &str,
|
||||
keep_alive: Option<&str>,
|
||||
) -> Result<api::LoadModelResponse, OllamaError> {
|
||||
let url = format!("{}/api/generate", self.base_url);
|
||||
|
||||
let keep_alive = keep_alive.ok_or(OllamaError::MissingKeepAlive)?;
|
||||
|
||||
self.parse_keep_alive(keep_alive)?;
|
||||
|
||||
let exists = self.model_exists(model).await?;
|
||||
if !exists {
|
||||
return Err(OllamaError::ModelNotFound(model.to_string()));
|
||||
}
|
||||
|
||||
let payload = json!({
|
||||
"model": model,
|
||||
"prompt": "",
|
||||
"keep_alive": keep_alive,
|
||||
"stream": false,
|
||||
});
|
||||
|
||||
let _res = self
|
||||
.client
|
||||
.post(url)
|
||||
.json(&payload)
|
||||
.send()
|
||||
.await?
|
||||
.text()
|
||||
.await?;
|
||||
|
||||
Ok(api::LoadModelResponse {
|
||||
model: model.to_string(),
|
||||
status: "loaded".to_string(),
|
||||
keep_alive: keep_alive.to_string(),
|
||||
})
|
||||
}
|
||||
|
||||
pub async fn unload_model(&self, model: &str) -> Result<api::UnloadModelResponse, OllamaError> {
|
||||
let url = format!("{}/api/generate", self.base_url);
|
||||
|
||||
let exists = self.model_exists(model).await?;
|
||||
if !exists {
|
||||
return Err(OllamaError::ModelNotFound(model.to_string()));
|
||||
}
|
||||
|
||||
let payload = json!({
|
||||
"model": model,
|
||||
"prompt": "",
|
||||
"keep_alive": "0",
|
||||
"stream": false,
|
||||
});
|
||||
|
||||
let _res = self
|
||||
.client
|
||||
.post(url)
|
||||
.json(&payload)
|
||||
.send()
|
||||
.await?
|
||||
.text()
|
||||
.await?;
|
||||
|
||||
Ok(api::UnloadModelResponse {
|
||||
model: model.to_string(),
|
||||
status: "unloaded".to_string(),
|
||||
})
|
||||
Ok(res)
|
||||
}
|
||||
|
||||
pub async fn completions(
|
||||
&self,
|
||||
body: &api::CompletionRequest,
|
||||
) -> Result<api::CompletionResponse, OllamaError> {
|
||||
body: &super::types::OllamaGenerateRequest,
|
||||
) -> Result<super::types::OllamaGenerateResponse, LlmError> {
|
||||
let url = format!("{}/api/generate", self.base_url);
|
||||
|
||||
let (prompt, model) = self.extract_completion_params(body)?;
|
||||
|
||||
if prompt.is_empty() {
|
||||
return Err(OllamaError::MissingPrompt);
|
||||
if body.prompt.is_empty() {
|
||||
return Err(LlmError::MissingPrompt);
|
||||
}
|
||||
|
||||
if model.is_empty() {
|
||||
return Err(OllamaError::MissingModel);
|
||||
if body.model.is_empty() {
|
||||
return Err(LlmError::MissingModel);
|
||||
}
|
||||
|
||||
let exists = self.model_exists(model).await?;
|
||||
let exists = self.model_exists(&body.model).await?;
|
||||
if !exists {
|
||||
return Err(OllamaError::ModelNotFound(model.to_string()));
|
||||
return Err(LlmError::ModelNotFound(body.model.clone()));
|
||||
}
|
||||
|
||||
let options = ollama::OllamaOptions::from(&body.base);
|
||||
let options = body.options.clone();
|
||||
|
||||
let payload = ollama::OllamaGenerateRequest {
|
||||
model,
|
||||
prompt,
|
||||
let payload = ollama::types::OllamaGenerateRequest {
|
||||
model: body.model.clone(),
|
||||
prompt: body.prompt.clone(),
|
||||
stream: false,
|
||||
keep_alive: body.keep_alive.clone(),
|
||||
options,
|
||||
};
|
||||
|
||||
dbg!(&payload);
|
||||
|
||||
let res = self
|
||||
.client
|
||||
.post(url)
|
||||
.json(&payload)
|
||||
.send()
|
||||
.await?
|
||||
.json::<ollama::OllamaGenerateResponse>()
|
||||
.json::<ollama::types::OllamaGenerateResponse>()
|
||||
.await?;
|
||||
|
||||
Ok(api::CompletionResponse::from(res))
|
||||
Ok(res)
|
||||
}
|
||||
|
||||
pub async fn completions_stream(
|
||||
&self,
|
||||
body: &api::CompletionRequest,
|
||||
) -> Result<ReceiverStream<Result<Event, OllamaError>>, OllamaError> {
|
||||
body: &super::types::OllamaGenerateRequest,
|
||||
) -> Result<ollama::types::OllamaGenerateResponseStream, LlmError> {
|
||||
let url = format!("{}/api/generate", self.base_url);
|
||||
|
||||
let (prompt, model) = self.extract_completion_params(body)?;
|
||||
|
||||
if prompt.is_empty() {
|
||||
return Err(OllamaError::MissingPrompt);
|
||||
if body.prompt.is_empty() {
|
||||
return Err(LlmError::MissingPrompt);
|
||||
}
|
||||
|
||||
if model.is_empty() {
|
||||
return Err(OllamaError::MissingModel);
|
||||
if body.model.is_empty() {
|
||||
return Err(LlmError::MissingModel);
|
||||
}
|
||||
|
||||
let exists = self.model_exists(model).await?;
|
||||
let exists = self.model_exists(&body.model.to_string()).await?;
|
||||
if !exists {
|
||||
return Err(OllamaError::ModelNotFound(model.to_string()));
|
||||
return Err(LlmError::ModelNotFound(body.model.to_string()));
|
||||
}
|
||||
|
||||
let options = ollama::OllamaOptions::from(&body.base);
|
||||
|
||||
let payload = ollama::OllamaGenerateRequest {
|
||||
model,
|
||||
prompt,
|
||||
let payload = ollama::types::OllamaGenerateRequest {
|
||||
model: body.model.clone(),
|
||||
prompt: body.prompt.clone(),
|
||||
stream: true,
|
||||
options,
|
||||
keep_alive: body.keep_alive.clone(),
|
||||
options: body.options.clone(),
|
||||
};
|
||||
|
||||
let mut byte_stream = self
|
||||
let byte_stream = self
|
||||
.client
|
||||
.post(url)
|
||||
.json(&payload)
|
||||
@@ -249,85 +156,80 @@ impl OllamaProvider {
|
||||
.await?
|
||||
.bytes_stream();
|
||||
|
||||
let (tx, rx) = tokio::sync::mpsc::channel(32);
|
||||
let stream = byte_stream.flat_map(|chunk_result| {
|
||||
let mut out: Vec<Result<super::types::OllamaGenerateStreamEvent, LlmError>> =
|
||||
Vec::new();
|
||||
|
||||
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;
|
||||
}
|
||||
};
|
||||
let chunk = match chunk_result {
|
||||
Ok(b) => b,
|
||||
Err(e) => {
|
||||
out.push(Err(LlmError::Http(e)));
|
||||
return futures::stream::iter(out);
|
||||
}
|
||||
};
|
||||
|
||||
// 🔥 IMPORTANT: typed deserialization
|
||||
let parsed: ollama::OllamaGenerateResponse = match serde_json::from_slice(&chunk) {
|
||||
Ok(v) => v,
|
||||
Err(_) => continue,
|
||||
};
|
||||
for line in chunk.split(|&b| b == b'\n') {
|
||||
if line.is_empty() {
|
||||
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 parsed: ollama::types::OllamaGenerateResponse =
|
||||
match serde_json::from_slice(line) {
|
||||
Ok(v) => v,
|
||||
Err(_) => continue,
|
||||
};
|
||||
|
||||
let _ = tx.send(Ok(Event::default().data(event_data))).await;
|
||||
if !parsed.response.is_empty() && !parsed.done {
|
||||
out.push(Ok(super::types::OllamaGenerateStreamEvent::Token(
|
||||
parsed.response.clone(),
|
||||
)));
|
||||
}
|
||||
|
||||
if parsed.done {
|
||||
let _ = tx.send(Ok(Event::default().data("[DONE]"))).await;
|
||||
break;
|
||||
out.push(Ok(super::types::OllamaGenerateStreamEvent::Final(parsed)));
|
||||
return futures::stream::iter(out);
|
||||
}
|
||||
}
|
||||
|
||||
futures::stream::iter(out)
|
||||
});
|
||||
|
||||
Ok(ReceiverStream::new(rx))
|
||||
Ok(Box::pin(stream))
|
||||
}
|
||||
|
||||
pub async fn chat_completions(
|
||||
&self,
|
||||
body: &api::ChatRequest,
|
||||
) -> Result<api::ChatCompletionResponse, OllamaError> {
|
||||
body: &super::types::OllamaChatRequest,
|
||||
) -> Result<super::types::OllamaChatResponse, LlmError> {
|
||||
let url = format!("{}/api/chat", self.base_url);
|
||||
|
||||
let (messages, model) = self.extract_chat_params(body)?;
|
||||
|
||||
if body.messages.is_empty() {
|
||||
return Err(OllamaError::MissingMessages);
|
||||
return Err(LlmError::MissingMessages);
|
||||
}
|
||||
|
||||
if !self.has_user_message(&body.messages) {
|
||||
return Err(OllamaError::MissingMessages);
|
||||
let ollama_messages: Vec<ollama::types::OllamaMessage> = body.messages.clone();
|
||||
|
||||
if !self.has_user_message(&ollama_messages) {
|
||||
return Err(LlmError::MissingMessages);
|
||||
}
|
||||
|
||||
if model.is_empty() {
|
||||
return Err(OllamaError::MissingModel);
|
||||
if body.model.is_empty() {
|
||||
return Err(LlmError::MissingModel);
|
||||
}
|
||||
|
||||
let exists = self.model_exists(model).await?;
|
||||
let exists = self.model_exists(&body.model).await?;
|
||||
if !exists {
|
||||
return Err(OllamaError::ModelNotFound(model.to_string()));
|
||||
return Err(LlmError::ModelNotFound(body.model.clone()));
|
||||
}
|
||||
|
||||
// let options = ollama::OllamaOptions::from(body);
|
||||
let options = ollama::OllamaOptions::from(&body.base);
|
||||
let options = body.options.clone();
|
||||
|
||||
let payload = ollama::OllamaChatRequest {
|
||||
model,
|
||||
messages,
|
||||
let payload = ollama::types::OllamaChatRequest {
|
||||
model: body.model.clone(),
|
||||
messages: ollama_messages,
|
||||
stream: false,
|
||||
options,
|
||||
keep_alive: body.keep_alive.clone(),
|
||||
};
|
||||
|
||||
let res = self
|
||||
@@ -336,43 +238,46 @@ impl OllamaProvider {
|
||||
.json(&payload)
|
||||
.send()
|
||||
.await?
|
||||
.json::<ollama::OllamaChatResponse>()
|
||||
.json::<ollama::types::OllamaChatResponse>()
|
||||
.await?;
|
||||
|
||||
Ok(api::ChatCompletionResponse::from(res))
|
||||
Ok(res)
|
||||
}
|
||||
|
||||
pub async fn chat_completions_stream(
|
||||
&self,
|
||||
body: &api::ChatRequest,
|
||||
) -> Result<ReceiverStream<Result<api::ChatCompletionChunk, OllamaError>>, OllamaError> {
|
||||
body: &super::types::OllamaChatRequest,
|
||||
) -> Result<ollama::types::OllamaChatResponseStream, LlmError> {
|
||||
let url = format!("{}/api/chat", self.base_url);
|
||||
|
||||
let (messages, model) = self.extract_chat_params(body)?;
|
||||
|
||||
if messages.is_empty() {
|
||||
return Err(OllamaError::MissingMessages);
|
||||
if body.messages.is_empty() {
|
||||
return Err(LlmError::MissingMessages);
|
||||
}
|
||||
|
||||
if model.is_empty() {
|
||||
return Err(OllamaError::MissingModel);
|
||||
if body.model.is_empty() {
|
||||
return Err(LlmError::MissingModel);
|
||||
}
|
||||
|
||||
let exists = self.model_exists(model).await?;
|
||||
let exists = self.model_exists(&body.model).await?;
|
||||
if !exists {
|
||||
return Err(OllamaError::ModelNotFound(model.to_string()));
|
||||
return Err(LlmError::ModelNotFound(body.model.clone()));
|
||||
}
|
||||
|
||||
let options = ollama::OllamaOptions::from(&body.base);
|
||||
let ollama_messages: Vec<ollama::types::OllamaMessage> = body.messages.clone();
|
||||
|
||||
let payload = ollama::OllamaChatRequest {
|
||||
model,
|
||||
messages,
|
||||
if !self.has_user_message(&ollama_messages) {
|
||||
return Err(LlmError::MissingMessages);
|
||||
}
|
||||
|
||||
let payload = ollama::types::OllamaChatRequest {
|
||||
model: body.model.clone(),
|
||||
messages: body.messages.clone().into_iter().collect(),
|
||||
stream: true,
|
||||
options,
|
||||
options: body.options.clone(),
|
||||
keep_alive: body.keep_alive.clone(),
|
||||
};
|
||||
|
||||
let mut byte_stream = self
|
||||
let byte_stream = self
|
||||
.client
|
||||
.post(url)
|
||||
.json(&payload)
|
||||
@@ -380,63 +285,42 @@ impl OllamaProvider {
|
||||
.await?
|
||||
.bytes_stream();
|
||||
|
||||
let (tx, rx) = tokio::sync::mpsc::channel(32);
|
||||
let stream = byte_stream.flat_map(|chunk_result| {
|
||||
let mut out: Vec<Result<super::types::OllamaChatStreamEvent, LlmError>> = Vec::new();
|
||||
|
||||
tokio::spawn(async move {
|
||||
let stream_id = format!("chatcmpl-{}", uuid::Uuid::new_v4());
|
||||
let chunk = match chunk_result {
|
||||
Ok(b) => b,
|
||||
Err(e) => {
|
||||
out.push(Err(LlmError::Http(e)));
|
||||
return futures::stream::iter(out);
|
||||
}
|
||||
};
|
||||
|
||||
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;
|
||||
}
|
||||
};
|
||||
for line in chunk.split(|&b| b == b'\n') {
|
||||
if line.is_empty() {
|
||||
continue;
|
||||
}
|
||||
|
||||
let parsed: ollama::OllamaChatResponse = match serde_json::from_slice(&chunk) {
|
||||
let parsed: ollama::types::OllamaChatResponse = match serde_json::from_slice(line) {
|
||||
Ok(v) => v,
|
||||
Err(_) => continue,
|
||||
};
|
||||
|
||||
let usage = if parsed.done {
|
||||
Some(api::Usage {
|
||||
prompt_tokens: parsed.prompt_eval_count.unwrap_or(0) as u32,
|
||||
completion_tokens: parsed.eval_count.unwrap_or(0) as u32,
|
||||
total_tokens: (parsed.prompt_eval_count.unwrap_or(0)
|
||||
+ parsed.eval_count.unwrap_or(0))
|
||||
as u32,
|
||||
})
|
||||
} else {
|
||||
None
|
||||
};
|
||||
|
||||
let event = api::ChatCompletionChunk {
|
||||
id: stream_id.clone(),
|
||||
object: "chat.completion.chunk".to_string(),
|
||||
choices: vec![api::ChatChunkChoice {
|
||||
index: 0,
|
||||
delta: api::Delta {
|
||||
role: Some(parsed.message.role),
|
||||
content: Some(parsed.message.content),
|
||||
},
|
||||
finish_reason: if parsed.done {
|
||||
Some(api::FinishReason::Stop)
|
||||
} else {
|
||||
None
|
||||
},
|
||||
}],
|
||||
usage,
|
||||
};
|
||||
|
||||
let _ = tx.send(Ok(event)).await;
|
||||
if !parsed.message.content.is_empty() && !parsed.done {
|
||||
out.push(Ok(super::types::OllamaChatStreamEvent::Token(
|
||||
parsed.message.content.clone(),
|
||||
)));
|
||||
}
|
||||
|
||||
if parsed.done {
|
||||
break;
|
||||
out.push(Ok(super::types::OllamaChatStreamEvent::Final(parsed)));
|
||||
return futures::stream::iter(out);
|
||||
}
|
||||
}
|
||||
|
||||
futures::stream::iter(out)
|
||||
});
|
||||
|
||||
Ok(ReceiverStream::new(rx))
|
||||
Ok(Box::pin(stream))
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,8 +1,7 @@
|
||||
// errors.rs
|
||||
use thiserror::Error;
|
||||
|
||||
#[derive(Debug, Error)]
|
||||
pub enum OllamaError {
|
||||
pub enum LlmError {
|
||||
#[error("prompt is required and cannot be empty")]
|
||||
MissingPrompt,
|
||||
|
||||
@@ -15,44 +14,10 @@ pub enum OllamaError {
|
||||
#[error("model '{0}' is not available — run `ollama pull {0}` first")]
|
||||
ModelNotFound(String),
|
||||
|
||||
#[error("keep_alive is required and cannot be empty")]
|
||||
MissingKeepAlive,
|
||||
|
||||
#[error(
|
||||
"invalid keep_alive format '{0}' — expected <number><unit> (e.g. 30s, 10m, 2h), a plain integer (seconds), or -1"
|
||||
)]
|
||||
InvalidKeepAlive(String),
|
||||
|
||||
// #[error(
|
||||
// "invalid keep_alive format '{0}' — expected <number><unit> (e.g. 30s, 10m, 2h), a plain integer (seconds), or -1"
|
||||
// )]
|
||||
// InvalidKeepAlive(String),
|
||||
#[error(transparent)]
|
||||
Http(#[from] reqwest::Error),
|
||||
}
|
||||
|
||||
pub fn into_http_response(e: OllamaError) -> (axum::http::StatusCode, String) {
|
||||
match e {
|
||||
OllamaError::MissingPrompt => (
|
||||
axum::http::StatusCode::BAD_REQUEST,
|
||||
"prompt is required and cannot be empty".to_string(),
|
||||
),
|
||||
OllamaError::MissingModel => (
|
||||
axum::http::StatusCode::BAD_REQUEST,
|
||||
"model is required and cannot be empty".to_string(),
|
||||
),
|
||||
OllamaError::ModelNotFound(m) => (
|
||||
axum::http::StatusCode::UNPROCESSABLE_ENTITY,
|
||||
format!("model '{m}' is not available — run `ollama pull {m}` first"),
|
||||
),
|
||||
OllamaError::MissingKeepAlive => (
|
||||
axum::http::StatusCode::BAD_REQUEST,
|
||||
"keep alive is required and cannot be empty".to_string(),
|
||||
),
|
||||
OllamaError::InvalidKeepAlive(v) => (
|
||||
axum::http::StatusCode::BAD_REQUEST,
|
||||
format!("invalid keep_alive '{v}'"),
|
||||
),
|
||||
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()),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,88 +0,0 @@
|
||||
use crate::dto::{api, ollama};
|
||||
|
||||
use chrono::Utc;
|
||||
use uuid::Uuid;
|
||||
|
||||
impl From<ollama::OllamaModel> for api::ModelInfo {
|
||||
fn from(m: ollama::OllamaModel) -> Self {
|
||||
Self {
|
||||
name: m.name,
|
||||
|
||||
family: m.details.as_ref().and_then(|d| d.family.clone()),
|
||||
parameter_size: m.details.as_ref().and_then(|d| d.parameter_size.clone()),
|
||||
quantization: m
|
||||
.details
|
||||
.as_ref()
|
||||
.and_then(|d| d.quantization_level.clone()),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl From<ollama::OllamaGenerateResponse> 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),
|
||||
},
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
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,
|
||||
}),
|
||||
|
||||
conversation_id: None,
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -1,3 +1,3 @@
|
||||
pub mod client;
|
||||
pub mod errors;
|
||||
pub mod mapper;
|
||||
pub mod types;
|
||||
|
||||
@@ -0,0 +1,137 @@
|
||||
use crate::providers::ollama::errors;
|
||||
|
||||
use futures::Stream;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::pin::Pin;
|
||||
|
||||
// ------ Models ------
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
pub struct OllamaModels {
|
||||
pub models: Vec<OllamaModel>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
pub struct OllamaModel {
|
||||
pub name: String,
|
||||
|
||||
pub details: Option<OllamaModelDetails>,
|
||||
|
||||
pub size: Option<u64>,
|
||||
pub digest: Option<String>,
|
||||
pub modified_at: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
pub struct OllamaModelDetails {
|
||||
pub family: Option<String>,
|
||||
pub parameter_size: Option<String>,
|
||||
pub quantization_level: Option<String>,
|
||||
}
|
||||
|
||||
// ------ Message ------
|
||||
|
||||
#[derive(Debug, Serialize, Deserialize, Clone)]
|
||||
pub struct OllamaMessage {
|
||||
pub role: OllamaRole,
|
||||
pub content: String,
|
||||
}
|
||||
|
||||
#[derive(Debug, Serialize, Deserialize, PartialEq, Eq, Clone)]
|
||||
#[serde(rename_all = "lowercase")]
|
||||
pub enum OllamaRole {
|
||||
System,
|
||||
User,
|
||||
Assistant,
|
||||
}
|
||||
|
||||
// ------ Shared ------
|
||||
|
||||
#[derive(Debug, Default, Clone, Serialize)]
|
||||
pub struct OllamaOptions {
|
||||
pub seed: Option<i64>,
|
||||
pub temperature: Option<f32>,
|
||||
pub top_p: Option<f32>,
|
||||
pub top_k: Option<u32>,
|
||||
pub stop: Option<Vec<String>>,
|
||||
pub num_ctx: Option<u32>,
|
||||
pub num_predict: Option<u32>,
|
||||
}
|
||||
|
||||
// ------ Completion ------
|
||||
|
||||
#[derive(Debug, Serialize)]
|
||||
pub struct OllamaGenerateRequest {
|
||||
pub model: String,
|
||||
pub prompt: String,
|
||||
pub stream: bool,
|
||||
pub keep_alive: String,
|
||||
pub options: Option<OllamaOptions>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
pub struct OllamaGenerateResponse {
|
||||
pub model: String,
|
||||
pub created_at: String,
|
||||
|
||||
pub response: String,
|
||||
|
||||
pub done: bool,
|
||||
pub done_reason: Option<String>,
|
||||
|
||||
pub total_duration: Option<u64>,
|
||||
pub load_duration: Option<u64>,
|
||||
|
||||
pub prompt_eval_count: Option<u32>,
|
||||
pub eval_count: Option<u32>,
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
pub enum OllamaGenerateStreamEvent {
|
||||
Token(String),
|
||||
Final(OllamaGenerateResponse),
|
||||
}
|
||||
|
||||
pub type OllamaGenerateResponseStream = Pin<
|
||||
Box<
|
||||
dyn Stream<Item = Result<super::types::OllamaGenerateStreamEvent, errors::LlmError>> + Send,
|
||||
>,
|
||||
>;
|
||||
|
||||
// ------ Chat ------
|
||||
|
||||
#[derive(Debug, Serialize)]
|
||||
pub struct OllamaChatRequest {
|
||||
pub model: String,
|
||||
pub messages: Vec<OllamaMessage>,
|
||||
pub stream: bool,
|
||||
pub keep_alive: String,
|
||||
pub options: Option<OllamaOptions>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
pub struct OllamaChatResponse {
|
||||
pub model: String,
|
||||
pub created_at: String,
|
||||
|
||||
pub message: OllamaMessage,
|
||||
|
||||
pub done: bool,
|
||||
pub done_reason: Option<String>,
|
||||
|
||||
pub total_duration: Option<u64>,
|
||||
pub load_duration: Option<u64>,
|
||||
|
||||
pub prompt_eval_count: Option<u32>,
|
||||
pub eval_count: Option<u32>,
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
pub enum OllamaChatStreamEvent {
|
||||
Token(String),
|
||||
Final(OllamaChatResponse),
|
||||
}
|
||||
|
||||
pub type OllamaChatResponseStream = Pin<
|
||||
Box<dyn Stream<Item = Result<super::types::OllamaChatStreamEvent, errors::LlmError>> + Send>,
|
||||
>;
|
||||
Reference in New Issue
Block a user