llm/providers/openai/
provider.rs1use async_openai::{Client, config::Config, types::chat::CreateChatCompletionRequest};
2use async_stream;
3use std::error::Error;
4use tokio_stream::StreamExt;
5use tracing::{debug, error};
6
7use super::{
8 mappers::{map_messages, map_tools},
9 streaming::process_completion_stream,
10};
11use crate::provider::error_stream;
12use crate::{Context, LlmError, LlmResponseStream, ProviderError, StreamingModelProvider};
13
14pub trait OpenAiChatProvider {
17 type Config: Config + Clone + 'static;
18
19 fn client(&self) -> &Client<Self::Config>;
20 fn model(&self) -> &str;
21 fn provider_name(&self) -> &str;
22}
23
24impl<T: OpenAiChatProvider + Send + Sync> StreamingModelProvider for T {
25 fn stream_response(&self, context: &Context) -> LlmResponseStream {
26 if let Err(error) = crate::provider::validate_reasoning(context, None) {
27 return crate::provider::error_stream(error);
28 }
29 let client = self.client().clone();
30 let model = self.model().to_string();
31 let messages = match map_messages(context.messages()) {
32 Ok(messages) => messages,
33 Err(e) => return error_stream(e),
34 };
35 let message_count = messages.len();
36 let tools = if context.tools().is_empty() {
37 None
38 } else {
39 match map_tools(context.tools(), None) {
40 Ok(t) => Some(t),
41 Err(e) => return error_stream(e),
42 }
43 };
44
45 Box::pin(async_stream::stream! {
46 debug!("Starting chat completion stream for model: {model}");
47
48 let req = CreateChatCompletionRequest {
49 model: model.clone(),
50 messages,
51 tools,
52 stream: Some(true),
53 ..Default::default()
54 };
55
56 debug!(
57 "Making request to Ollama API with model: {model} and {message_count} messages"
58 );
59
60 let stream = match client.chat().create_stream(req).await {
61 Ok(stream) => {
62 debug!("Successfully created stream from Ollama API");
63 stream
64 }
65 Err(e) => {
66 error!("Failed to create stream from Ollama API: {:?}", e);
67
68 if let Some(reqwest_err) =
70 e.source().and_then(|s| s.downcast_ref::<reqwest::Error>())
71 {
72 if let Some(url) = reqwest_err.url() {
73 error!("Request URL was: {url}");
74 }
75 if let Some(status) = reqwest_err.status() {
76 error!("HTTP status: {status}");
77 }
78 }
79
80 yield Err(LlmError::from(e));
81 return;
82 }
83 };
84
85 let stream = stream.map(|result| {
86 result.map_err(|e| LlmError::from(ProviderError::stream_interrupted(e.to_string())))
87 });
88
89 let mut shared_stream = Box::pin(process_completion_stream(stream));
90 while let Some(result) = shared_stream.next().await {
91 yield result;
92 }
93 })
94 }
95
96 fn context_window(&self) -> Option<u32> {
97 None
98 }
99
100 fn display_name(&self) -> String {
101 let model = self.model();
102 if model.is_empty() { self.provider_name().to_string() } else { format!("{} ({model})", self.provider_name()) }
103 }
104}
105
106#[cfg(test)]
107mod tests {
108 use futures::StreamExt;
109
110 use super::*;
111 use crate::providers::local::ollama::OllamaProvider;
112 use crate::providers::test_capture_server::CaptureServer;
113 use crate::{ChatMessage, LlmResponse, Result};
114
115 #[tokio::test]
116 async fn local_chat_providers_omit_cache_metadata() {
117 let mut server = CaptureServer::start_chat_completions().await;
118 let provider = OllamaProvider::new("test-model", &server.base_url);
119 let mut context = Context::new(vec![ChatMessage::user("Hello")], vec![]);
120 context.set_prompt_cache_key(Some("prefix-abc".to_string()));
121 context.set_session_affinity_key(Some("conversation-abc".to_string()));
122
123 let responses = provider.stream_response(&context).collect::<Vec<_>>().await;
124 let captured = server.captured().await;
125
126 assert!(responses.iter().all(Result::is_ok), "{responses:?}");
127 assert!(responses.iter().any(|response| matches!(response, Ok(LlmResponse::Done { .. }))));
128 assert_eq!(captured.path, "/v1/chat/completions");
129 assert!(captured.body.get("prompt_cache_key").is_none());
130 assert!(captured.body.get("session_id").is_none());
131 assert!(captured.body.get("user").is_none());
132 }
133}