rig_core/providers/ollama/
streaming.rs1use serde::Deserialize;
17use serde_json::{Map, Value};
18
19use crate::completion::{FinishReason, Usage};
20use crate::error::ProviderError;
21use crate::json_utils::Lenient;
22use crate::message::{CallId, ToolName};
23use crate::observe::ObservedError;
24use crate::operation::{Block, Completion, Finish};
25use crate::providers::internal::wire;
26use crate::wire::{
27 AdapterEvent, AdapterUsage, AdapterVerdict, Decoder, Flow, ObservationSink, Out, WireEvent,
28 WireFrame,
29};
30
31const RECORD_KEYS: &[&str] = &["message", "done", "error"];
34
35const THINK_OPEN: &str = "<think>";
37const THINK_CLOSE: &str = "</think>";
39
40#[derive(Debug, Default, Deserialize)]
42#[serde(transparent)]
43pub struct ChatRecord(pub Map<String, Value>);
44
45#[derive(Debug, Default)]
62pub struct ChatDecoder {
63 split: Split,
65 called: bool,
67 model: Option<String>,
69}
70
71#[derive(Debug)]
73enum Split {
74 Opening(String),
77 Inside { held: String, started: bool },
81 Text { trim: bool },
84}
85
86impl Default for Split {
87 fn default() -> Self {
88 Self::Opening(String::new())
89 }
90}
91
92fn partial_suffix(text: &str, tag: &str) -> usize {
95 (1..tag.len())
96 .rev()
97 .find(|&len| text.ends_with(&tag[..len]))
98 .unwrap_or(0)
99}
100
101impl<'id> Decoder<'id, Completion> for ChatDecoder {
102 type Event = ChatRecord;
103
104 fn classify(&self, frame: WireFrame) -> WireEvent<ChatRecord> {
105 wire::classify_marker_keyed_frame(&frame.as_str(), RECORD_KEYS)
106 }
107
108 fn decode(
111 &mut self,
112 ChatRecord(fields): ChatRecord,
113 mut out: Out<'id, Completion>,
114 ) -> Result<Flow, ProviderError> {
115 let record = Value::Object(fields);
116 if let Some(error) = record.get("error").filter(|error| !error.is_null()) {
117 let body = serde_json::json!({ "error": error }).to_string();
118 return Err(ProviderError::from_provider_body(body));
119 }
120 if let Some(model) = record.str("model").filter(|model| !model.is_empty()) {
121 tracing::Span::current().record("gen_ai.response.model", model);
122 self.model = Some(model.to_owned());
123 }
124 let message = record.get("message").unwrap_or(&Value::Null);
125 if let Some(thinking) = message
126 .str("thinking")
127 .filter(|thinking| !thinking.is_empty())
128 {
129 self.release(&mut out)?;
131 reason(thinking, &mut out)?;
132 }
133 if let Some(content) = message.str("content").filter(|content| !content.is_empty()) {
134 self.content(content, &mut out)?;
135 }
136 for call in message.arr("tool_calls") {
137 self.call(call, &mut out)?;
138 }
139 if record.bool("done") != Some(true) {
140 return Ok(Flow::More);
141 }
142 self.release(&mut out)?;
143 out.end_run()?;
144 let reason = record.str("done_reason").map(|reason| match reason {
145 "stop" if self.called => FinishReason::ToolCalls,
146 "stop" => FinishReason::Stop,
147 "length" => FinishReason::Length,
148 other => FinishReason::Other(other.to_owned()),
149 });
150 let (input, output) = (record.u64("prompt_eval_count"), record.u64("eval_count"));
151 let usage = Usage {
152 input_tokens: input,
153 output_tokens: output,
154 total_tokens: input.zip(output).map(|(input, output)| input + output),
155 cached_input_tokens: record.u64("prompt_eval_cached_count"),
156 ..Usage::default()
157 };
158 Ok(out.end(Finish {
159 usage,
160 reason,
161 model: self.model.take(),
162 ..Finish::default()
163 }))
164 }
165}
166
167impl ChatDecoder {
168 fn content(
170 &mut self,
171 content: &str,
172 out: &mut Out<'_, Completion>,
173 ) -> Result<(), ProviderError> {
174 match std::mem::replace(&mut self.split, Split::Text { trim: false }) {
175 Split::Text { trim } => {
176 let text = if trim { content.trim_start() } else { content };
177 self.split = Split::Text {
178 trim: trim && text.is_empty(),
179 };
180 if !text.is_empty() {
181 out.run(Block::Text, text)?;
182 }
183 Ok(())
184 }
185 Split::Opening(mut held) => {
186 held.push_str(content);
187 let trimmed = held.trim_start();
188 if let Some(rest) = trimmed.strip_prefix(THINK_OPEN) {
189 let rest = rest.to_owned();
190 self.split = Split::Inside {
191 held: String::new(),
192 started: false,
193 };
194 self.content(&rest, out)
195 } else if THINK_OPEN.starts_with(trimmed) {
196 self.split = Split::Opening(held);
197 Ok(())
198 } else {
199 out.run(Block::Text, &held)?;
200 Ok(())
201 }
202 }
203 Split::Inside {
204 mut held,
205 mut started,
206 } => {
207 held.push_str(content);
208 if let Some((reasoning, rest)) = held.split_once(THINK_CLOSE) {
209 write_reasoning(reasoning.trim_end(), started, out)?;
210 self.split = Split::Text { trim: true };
211 return self.content(rest, out);
212 }
213 let cut = held.len() - partial_suffix(&held, THINK_CLOSE);
214 let cut = held[..cut].trim_end().len();
215 started |= write_reasoning(&held[..cut], started, out)?;
216 self.split = Split::Inside {
217 held: held.split_off(cut),
218 started,
219 };
220 Ok(())
221 }
222 }
223 }
224
225 fn release(&mut self, out: &mut Out<'_, Completion>) -> Result<(), ProviderError> {
228 match std::mem::replace(&mut self.split, Split::Text { trim: false }) {
229 Split::Opening(held) if !held.is_empty() => {
230 out.run(Block::Text, &held)?;
231 }
232 Split::Inside { held, started } => {
233 write_reasoning(held.trim_end(), started, out)?;
234 }
235 Split::Opening(_) | Split::Text { .. } => {}
236 }
237 Ok(())
238 }
239
240 fn call(&mut self, call: &Value, out: &mut Out<'_, Completion>) -> Result<(), ProviderError> {
243 let Ok(name) = ToolName::new(
244 call.at("/function/name")
245 .and_then(Value::as_str)
246 .unwrap_or_default(),
247 ) else {
248 tracing::warn!("Ollama sent a tool call without a name; nothing can answer it");
249 return Ok(());
250 };
251 let arguments = match call.at("/function/arguments") {
252 None | Some(Value::Null) => "{}".to_owned(),
253 Some(Value::String(arguments)) => arguments.clone(),
254 Some(arguments) => arguments.to_string(),
255 };
256 self.release(out)?;
257 out.end_run()?;
258 self.called = true;
259 let id = CallId::from_wire(call.str("id").unwrap_or_default());
260 let index = out.fresh_index();
261 out.open(index, Block::Call { id, name }, call.clone())?;
262 out.push(index, &arguments)?;
263 out.finish(index)
264 }
265
266 pub(crate) fn project(payload: &[u8], sink: &mut ObservationSink<'_>) {
270 let Ok(record) = serde_json::from_slice::<Value>(payload) else {
271 return;
272 };
273 let input = record.u64("prompt_eval_count");
274 let output = record.u64("eval_count");
275 if input.is_some() || output.is_some() {
276 sink.emit(AdapterEvent::Usage {
277 usage: AdapterUsage {
278 input_tokens: input,
279 output_tokens: output,
280 total_tokens: input.zip(output).map(|(input, output)| input + output),
281 cached_input_tokens: record.u64("prompt_eval_cached_count"),
282 reasoning_tokens: None,
283 tool_input_tokens: None,
284 },
285 });
286 }
287 let verdict = match record.str("done_reason") {
288 Some(reason) => AdapterVerdict {
289 finish_reason: Some(sink.scrub(reason)),
290 block_reason: None,
291 detail: None,
292 model: record.str("model").map(|model| sink.scrub(model)),
293 },
294 None => AdapterVerdict::default(),
295 };
296 sink.provider(verdict, None);
297 if let Some(message) = record.str("error") {
298 ObservedError {
299 code: None,
300 kind: None,
301 message: Some(message.to_owned()),
302 }
303 .emit(sink);
304 }
305 }
306}
307
308fn write_reasoning(
311 text: &str,
312 started: bool,
313 out: &mut Out<'_, Completion>,
314) -> Result<bool, ProviderError> {
315 let text = if started { text } else { text.trim_start() };
316 if text.is_empty() {
317 return Ok(false);
318 }
319 reason(text, out)?;
320 Ok(true)
321}
322
323fn reason(text: &str, out: &mut Out<'_, Completion>) -> Result<(), ProviderError> {
326 let index = out.run(Block::Reasoning { redacted: false }, text)?;
327 out.edit(index, |item| match item {
328 Value::Object(fields) => {
329 if let Some(Value::String(thinking)) = fields.get_mut("thinking") {
330 thinking.push_str(text);
331 }
332 }
333 _ => *item = serde_json::json!({ "thinking": text }),
334 })
335}
336
337pub(crate) mod document;
338
339#[cfg(test)]
340mod tests;