feat: add proper typing for api
This commit is contained in:
@@ -0,0 +1,409 @@
|
||||
use crate::dto::{api, ollama};
|
||||
use crate::errors::OllamaError;
|
||||
use axum::response::sse::Event;
|
||||
use futures::StreamExt;
|
||||
use reqwest::Client;
|
||||
use serde_json::{Value, json};
|
||||
use tokio_stream::wrappers::ReceiverStream;
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct OllamaProvider {
|
||||
pub client: Client,
|
||||
pub base_url: String,
|
||||
}
|
||||
|
||||
impl OllamaProvider {
|
||||
pub fn new(base_url: impl Into<String>) -> Self {
|
||||
Self {
|
||||
client: Client::new(),
|
||||
base_url: base_url.into(),
|
||||
}
|
||||
}
|
||||
|
||||
// ── private helpers ──────────────────────────────────────────────────────
|
||||
|
||||
async fn model_exists(&self, model: &str) -> Result<bool, OllamaError> {
|
||||
let url = format!("{}/api/tags", self.base_url);
|
||||
|
||||
let res = self
|
||||
.client
|
||||
.get(url)
|
||||
.send()
|
||||
.await?
|
||||
.json::<ollama::OllamaModels>()
|
||||
.await?;
|
||||
|
||||
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"),
|
||||
// })
|
||||
// }
|
||||
|
||||
// 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);
|
||||
|
||||
// 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))
|
||||
// }
|
||||
|
||||
// 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 ─────────────────────────────────────────────────────
|
||||
|
||||
pub async fn list_models(&self) -> Result<api::ModelsResponse, OllamaError> {
|
||||
let url = format!("{}/api/tags", self.base_url);
|
||||
|
||||
let res = self
|
||||
.client
|
||||
.get(url)
|
||||
.send()
|
||||
.await?
|
||||
.json::<ollama::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: &str,
|
||||
) -> Result<api::LoadModelResponse, 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": keep_alive,
|
||||
"stream": false,
|
||||
});
|
||||
|
||||
let _res = self
|
||||
.client
|
||||
.post(url)
|
||||
.json(&payload)
|
||||
.send()
|
||||
.await?
|
||||
.json::<Value>()
|
||||
.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?
|
||||
.json::<Value>()
|
||||
.await?;
|
||||
|
||||
Ok(api::UnloadModelResponse {
|
||||
model: model.to_string(),
|
||||
status: "unloaded".to_string(),
|
||||
})
|
||||
}
|
||||
|
||||
// pub async fn completions(&self, body: Value) -> Result<Value, OllamaError> {
|
||||
// let (prompt, model) = self.extract_completion_params(&body)?;
|
||||
// self.validate_model(model).await?;
|
||||
|
||||
// let payload = json!({
|
||||
// "model": model,
|
||||
// "prompt": prompt,
|
||||
// "stream": false,
|
||||
// "options": Self::build_options(&body),
|
||||
// });
|
||||
|
||||
// let res = self
|
||||
// .client
|
||||
// .post(format!("{}/api/generate", self.base_url))
|
||||
// .json(&payload)
|
||||
// .send()
|
||||
// .await?
|
||||
// .json::<Value>()
|
||||
// .await?;
|
||||
|
||||
// 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",
|
||||
// "object": "text_completion",
|
||||
// "choices": [{ "text": token, "index": 0, "finish_reason": 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> {
|
||||
// let (model, messages) = self.extract_chat_params(&body)?;
|
||||
// self.validate_model(model).await?;
|
||||
|
||||
// let payload = json!({
|
||||
// "model": model,
|
||||
// "messages": messages,
|
||||
// "stream": false,
|
||||
// "options": Self::build_options(&body),
|
||||
// });
|
||||
|
||||
// let res = self
|
||||
// .client
|
||||
// .post(format!("{}/api/chat", self.base_url))
|
||||
// .json(&payload)
|
||||
// .send()
|
||||
// .await?
|
||||
// .json::<Value>()
|
||||
// .await?;
|
||||
|
||||
// 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",
|
||||
// "object": "chat.completion.chunk",
|
||||
// "choices": [{
|
||||
// "index": 0,
|
||||
// "delta": {
|
||||
// "role": json.get("message").and_then(|m| m.get("role")),
|
||||
// "content": json.get("message").and_then(|m| m.get("content")),
|
||||
// },
|
||||
// "finish_reason": if done { json!("stop") } else { json!(null) },
|
||||
// }],
|
||||
// }))
|
||||
// .unwrap_or_default();
|
||||
|
||||
// let _ = tx.send(Ok(Event::default().data(event_data))).await;
|
||||
|
||||
// if done {
|
||||
// let _ = tx.send(Ok(Event::default().data("[DONE]"))).await;
|
||||
// break;
|
||||
// }
|
||||
// }
|
||||
// }
|
||||
// });
|
||||
|
||||
// Ok(ReceiverStream::new(rx))
|
||||
// }
|
||||
}
|
||||
@@ -0,0 +1,16 @@
|
||||
use crate::dto::{api, ollama};
|
||||
|
||||
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()),
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,2 @@
|
||||
pub mod client;
|
||||
pub mod mapper;
|
||||
Reference in New Issue
Block a user