Skip to main content

rig_core/providers/ollama/
wire.rs

1//! Ollama daemon configuration and chat, embedding, and model-listing wires.
2//!
3//! ```
4//! use rig_core::providers::ollama::OllamaConfig;
5//! let ollama = OllamaConfig::new().client();
6//! assert_eq!(ollama.completion("qwen3").wire.model, "qwen3");
7//! ```
8
9use 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
28/// The environment variable overriding the daemon's address.
29const BASE_URL_ENV: &str = "OLLAMA_API_BASE_URL";
30
31/// The environment variable carrying the bearer token a proxied daemon
32/// requires. A local daemon needs none.
33const API_KEY_ENV: &str = "OLLAMA_API_KEY";
34
35/// The most texts `POST /api/embed` accepts in one call.
36const MAX_DOCUMENTS: usize = 1024;
37
38/// The settings of an Ollama daemon: serializable, and the credential is
39/// never serialized. [`connect`](Self::connect) puts it on a transport as an
40/// [`Ollama`](super::Ollama) client.
41#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
42pub struct OllamaConfig {
43    /// The daemon's address.
44    pub base_url: String,
45    /// Bearer token for a secured daemon. Empty by default; no credential is
46    /// sent when empty.
47    pub api_key: Secret,
48}
49
50impl Default for OllamaConfig {
51    fn default() -> Self {
52        Self::new()
53    }
54}
55
56impl OllamaConfig {
57    /// The local daemon, unauthenticated.
58    pub fn new() -> Self {
59        Self {
60            base_url: OLLAMA_API_BASE_URL.to_owned(),
61            api_key: Secret::default(),
62        }
63    }
64
65    /// The daemon `OLLAMA_API_BASE_URL` names, with the token
66    /// `OLLAMA_API_KEY` carries. Both are optional: an unset base URL is the
67    /// local daemon and an unset key is no credential.
68    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    /// Point the wires at another daemon.
80    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    /// Authenticate against a proxied daemon.
86    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    /// The chat wire for `model`.
92    pub(crate) fn completion(&self, model: impl Into<String>) -> Chat {
93        Chat {
94            provider: self.clone(),
95            model: model.into(),
96        }
97    }
98
99    /// Build an embedding wire reporting the supplied width, known model width,
100    /// or zero if unknown. The width is metadata and is not sent to the daemon.
101    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    /// The model-listing wire.
114    pub(crate) fn models(&self) -> Models {
115        Models {
116            provider: self.clone(),
117        }
118    }
119
120    /// One request to `path`, with the credential only when there is one.
121    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/// The chat wire: `POST /api/chat`, NDJSON when streamed.
138#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
139pub struct Chat {
140    /// The daemon this wire speaks to.
141    pub provider: OllamaConfig,
142    /// The model to address.
143    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        // Both modes decode the same record shape; streaming needs NDJSON framing.
170        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/// The embedding wire: `POST /api/embed`.
185#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
186pub struct Embeddings {
187    /// The daemon this wire speaks to.
188    pub provider: OllamaConfig,
189    /// The model to address.
190    pub model: String,
191    /// The width this wire reports, from the caller or the model's published
192    /// dimensions. `0` means neither named one.
193    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
222/// Decodes one `/api/embed` reply.
223pub 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        // Ollama counts the prompt it embedded and nothing else: every token
241        // of an embedding is input.
242        let usage = crate::completion::Usage {
243            input_tokens: reply.prompt_eval_count,
244            total_tokens: reply.prompt_eval_count,
245            ..Default::default()
246        };
247        // The vectors only; the operation's fold pairs them with the texts
248        // that were sent, which `/api/embed` does not echo back.
249        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/// The model-listing wire: `GET /api/tags`, unpaged.
266#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
267pub struct Models {
268    /// The daemon this wire speaks to.
269    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
295/// Decodes `GET /api/tags`. The daemon answers with every installed model at
296/// once, so there is no cursor to follow.
297pub 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;