rig_core/providers/ollama/
wire.rs1use crate::client::env::{self, EnvError};
10use crate::completion::CompletionRequest;
11use crate::embeddings::Embedding as Vector;
12use crate::error::EncodeError;
13use crate::error::ProviderError;
14use crate::model::{ModelInfo, ModelList};
15use crate::operation::{Completion, Embedding, ModelListing, ModelPage};
16use crate::wire::Flow;
17use crate::wire::{
18 Body, Capabilities, Decoder, Descriptor, Encoded, Framing, Mode, Out, Secret, Wire, WireEvent,
19 WireFrame,
20};
21use serde::{Deserialize, Serialize};
22
23use super::{
24 EmbeddingResponse as OllamaEmbeddingResponse, ListModelsResponse, OLLAMA_API_BASE_URL,
25 OllamaCompletionRequest, OllamaDecoder, PROVIDER_NAME, model_dimensions_from_identifier,
26};
27
28const BASE_URL_ENV: &str = "OLLAMA_API_BASE_URL";
30
31const API_KEY_ENV: &str = "OLLAMA_API_KEY";
34
35const MAX_DOCUMENTS: usize = 1024;
37
38#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
42pub struct OllamaConfig {
43 pub base_url: String,
45 pub api_key: Secret,
48}
49
50impl Default for OllamaConfig {
51 fn default() -> Self {
52 Self::new()
53 }
54}
55
56impl OllamaConfig {
57 pub fn new() -> Self {
59 Self {
60 base_url: OLLAMA_API_BASE_URL.to_owned(),
61 api_key: Secret::default(),
62 }
63 }
64
65 pub fn from_env() -> Result<Self, EnvError> {
69 let mut provider = Self::new();
70 if let Some(base_url) = env::optional(BASE_URL_ENV)? {
71 provider = provider.with_base_url(base_url);
72 }
73 if let Some(api_key) = env::optional(API_KEY_ENV)? {
74 provider.api_key = api_key.into();
75 }
76 Ok(provider)
77 }
78
79 pub fn with_base_url(mut self, base_url: impl AsRef<str>) -> Self {
81 self.base_url = base_url.as_ref().trim_end_matches('/').to_owned();
82 self
83 }
84
85 pub fn with_api_key(mut self, api_key: impl Into<Secret>) -> Self {
87 self.api_key = api_key.into();
88 self
89 }
90
91 pub(crate) fn completion(&self, model: impl Into<String>) -> Chat {
93 Chat {
94 provider: self.clone(),
95 model: model.into(),
96 }
97 }
98
99 pub(crate) fn embedding(&self, model: impl Into<String>, ndims: Option<usize>) -> Embeddings {
102 let model = model.into();
103 let ndims = ndims
104 .or_else(|| model_dimensions_from_identifier(&model))
105 .unwrap_or_default();
106 Embeddings {
107 provider: self.clone(),
108 model,
109 ndims,
110 }
111 }
112
113 pub(crate) fn models(&self) -> Models {
115 Models {
116 provider: self.clone(),
117 }
118 }
119
120 fn request(&self, method: http::Method, path: &str) -> http::request::Builder {
122 let builder = http::Request::builder()
123 .method(method)
124 .uri(format!("{}{path}", self.base_url))
125 .header(http::header::CONTENT_TYPE, "application/json");
126 if self.api_key.is_empty() {
127 builder
128 } else {
129 builder.header(
130 http::header::AUTHORIZATION,
131 format!("Bearer {}", self.api_key.expose()),
132 )
133 }
134 }
135}
136
137#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
139pub struct Chat {
140 pub provider: OllamaConfig,
142 pub model: String,
144}
145
146impl Wire for Chat {
147 type Op = Completion;
148 type Payload = crate::wire::Encoded;
149 type Frame = crate::wire::WireFrame;
150 type Decoder<'id> = OllamaDecoder<'id>;
151
152 fn describe(&self) -> Descriptor<'_> {
153 Descriptor::new(PROVIDER_NAME).model(self.model.as_str())
154 }
155
156 fn encode(&self, request: CompletionRequest, mode: Mode) -> Result<Encoded, EncodeError> {
157 let request = request.replayable_to(&[super::ISSUER])?;
158 let mut body = OllamaCompletionRequest::try_from((self.model.as_str(), request))?;
159 body.stream = mode == Mode::Streaming;
160 crate::providers::internal::trace_json(
161 crate::providers::internal::LogTarget::Completions,
162 "Ollama completion request",
163 &body,
164 );
165 let request = self
166 .provider
167 .request(http::Method::POST, "/api/chat")
168 .body(Body::Bytes(serde_json::to_vec(&body)?))?;
169 Ok(Encoded::new(
171 request,
172 match mode {
173 Mode::Unary => Framing::Whole,
174 Mode::Streaming => Framing::Ndjson,
175 },
176 ))
177 }
178
179 fn decoder<'id>(&self) -> Self::Decoder<'id> {
180 OllamaDecoder::default()
181 }
182}
183
184#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
186pub struct Embeddings {
187 pub provider: OllamaConfig,
189 pub model: String,
191 pub ndims: usize,
194}
195
196impl Wire for Embeddings {
197 type Op = Embedding;
198 type Payload = crate::wire::Encoded;
199 type Frame = crate::wire::WireFrame;
200 type Decoder<'id> = EmbeddingsDecoder;
201
202 fn describe(&self) -> Descriptor<'_> {
203 Descriptor::new(PROVIDER_NAME)
204 .model(self.model.as_str())
205 .capabilities(Capabilities::embedding(MAX_DOCUMENTS, self.ndims))
206 }
207
208 fn encode(&self, texts: Vec<String>, _mode: Mode) -> Result<Encoded, EncodeError> {
209 let body = serde_json::json!({ "model": self.model, "input": texts });
210 let request = self
211 .provider
212 .request(http::Method::POST, "/api/embed")
213 .body(Body::Bytes(serde_json::to_vec(&body)?))?;
214 Ok(Encoded::new(request, Framing::Whole))
215 }
216
217 fn decoder<'id>(&self) -> Self::Decoder<'id> {
218 EmbeddingsDecoder
219 }
220}
221
222pub struct EmbeddingsDecoder;
224
225impl<'id> Decoder<'id, Embedding> for EmbeddingsDecoder {
226 type Event = OllamaEmbeddingResponse;
227
228 fn classify(&self, frame: WireFrame) -> WireEvent<Self::Event> {
229 crate::providers::internal::wire::classify_marker_keyed_frame(
230 &frame.as_str(),
231 &["embeddings"],
232 )
233 }
234
235 fn decode(
236 &mut self,
237 reply: Self::Event,
238 out: Out<'id, Embedding>,
239 ) -> Result<Flow, ProviderError> {
240 let usage = crate::completion::Usage {
243 input_tokens: reply.prompt_eval_count,
244 total_tokens: reply.prompt_eval_count,
245 ..Default::default()
246 };
247 let vectors = reply
250 .embeddings
251 .into_iter()
252 .map(|vec| Vector {
253 document: String::new(),
254 vec,
255 })
256 .collect();
257 Ok(out.end(crate::embeddings::EmbeddingResponse {
258 model: Some(reply.model),
259 usage,
260 ..crate::embeddings::EmbeddingResponse::new(vectors)
261 }))
262 }
263}
264
265#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
267pub struct Models {
268 pub provider: OllamaConfig,
270}
271
272impl Wire for Models {
273 type Op = ModelListing;
274 type Payload = crate::wire::Encoded;
275 type Frame = crate::wire::WireFrame;
276 type Decoder<'id> = ModelsDecoder;
277
278 fn describe(&self) -> Descriptor<'_> {
279 Descriptor::new(PROVIDER_NAME)
280 }
281
282 fn encode(&self, _cursor: Option<String>, _mode: Mode) -> Result<Encoded, EncodeError> {
283 let request = self
284 .provider
285 .request(http::Method::GET, "/api/tags")
286 .body(Body::empty())?;
287 Ok(Encoded::new(request, Framing::Whole))
288 }
289
290 fn decoder<'id>(&self) -> Self::Decoder<'id> {
291 ModelsDecoder
292 }
293}
294
295pub struct ModelsDecoder;
298
299impl<'id> Decoder<'id, ModelListing> for ModelsDecoder {
300 type Event = ListModelsResponse;
301
302 fn classify(&self, frame: WireFrame) -> WireEvent<Self::Event> {
303 crate::providers::internal::wire::classify_marker_keyed_frame(&frame.as_str(), &["models"])
304 }
305
306 fn decode(
307 &mut self,
308 reply: Self::Event,
309 out: Out<'id, ModelListing>,
310 ) -> Result<Flow, ProviderError> {
311 Ok(out.end(ModelPage {
312 models: ModelList::new(reply.models.into_iter().map(ModelInfo::from).collect()),
313 next: None,
314 }))
315 }
316}
317
318#[cfg(test)]
319mod tests;