reafctor: all code without stream

This commit is contained in:
2026-06-02 18:30:39 +02:00
parent 974be437af
commit 28352d8bdb
90 changed files with 3581 additions and 2434 deletions
+23
View File
@@ -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>,
}
+40
View File
@@ -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),
}
+59
View File
@@ -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
}
+4
View File
@@ -0,0 +1,4 @@
pub mod claims;
pub mod errors;
pub mod jwks;
pub mod validator;
+64
View File
@@ -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
View File
@@ -1 +1,2 @@
pub mod keycloak;
pub mod ollama;
+149 -265
View File
@@ -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))
}
}
+5 -40
View File
@@ -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()),
}
}
-88
View File
@@ -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 -1
View File
@@ -1,3 +1,3 @@
pub mod client;
pub mod errors;
pub mod mapper;
pub mod types;
+137
View File
@@ -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>,
>;