Skip to main content

llm/providers/openai/
provider.rs

1use 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, StreamingModelProvider};
13
14/// A Provider that's compatible with `OpenAI`'s chat completion API
15/// Other providers (e.g. Ollama, Llama.cpp etc) that are "`OpenAI` compatible" should implement this trait
16pub 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        let client = self.client().clone();
27        let model = self.model().to_string();
28        let messages = match map_messages(context.messages()) {
29            Ok(messages) => messages,
30            Err(e) => return error_stream(e),
31        };
32        let message_count = messages.len();
33        let tools = if context.tools().is_empty() {
34            None
35        } else {
36            match map_tools(context.tools(), None) {
37                Ok(t) => Some(t),
38                Err(e) => return error_stream(e),
39            }
40        };
41
42        Box::pin(async_stream::stream! {
43            debug!("Starting chat completion stream for model: {model}");
44
45            let req = CreateChatCompletionRequest {
46                model: model.clone(),
47                messages,
48                tools,
49                stream: Some(true),
50                ..Default::default()
51            };
52
53            debug!(
54                "Making request to Ollama API with model: {model} and {message_count} messages"
55            );
56
57            let stream = match client.chat().create_stream(req).await {
58                Ok(stream) => {
59                    debug!("Successfully created stream from Ollama API");
60                    stream
61                }
62                Err(e) => {
63                    error!("Failed to create stream from Ollama API: {:?}", e);
64
65                    // Check if it's a reqwest error with more details
66                    if let Some(reqwest_err) =
67                        e.source().and_then(|s| s.downcast_ref::<reqwest::Error>())
68                    {
69                        if let Some(url) = reqwest_err.url() {
70                            error!("Request URL was: {url}");
71                        }
72                        if let Some(status) = reqwest_err.status() {
73                            error!("HTTP status: {status}");
74                        }
75                    }
76
77                    yield Err(LlmError::ApiRequest(e.to_string()));
78                    return;
79                }
80            };
81
82            let stream = stream.map(|result| {
83                result.map_err(|e| LlmError::ApiError(e.to_string()))
84            });
85
86            let mut shared_stream = Box::pin(process_completion_stream(stream));
87            while let Some(result) = shared_stream.next().await {
88                yield result;
89            }
90        })
91    }
92
93    fn context_window(&self) -> Option<u32> {
94        None
95    }
96
97    fn display_name(&self) -> String {
98        let model = self.model();
99        if model.is_empty() { self.provider_name().to_string() } else { format!("{} ({model})", self.provider_name()) }
100    }
101}
102
103#[cfg(test)]
104mod tests {
105    use futures::StreamExt;
106
107    use super::*;
108    use crate::providers::local::ollama::OllamaProvider;
109    use crate::providers::test_capture_server::CaptureServer;
110    use crate::{ChatMessage, LlmResponse, Result};
111
112    #[tokio::test]
113    async fn local_chat_providers_omit_cache_metadata() {
114        let mut server = CaptureServer::start_chat_completions().await;
115        let provider = OllamaProvider::new("test-model", &server.base_url);
116        let mut context = Context::new(vec![ChatMessage::user("Hello")], vec![]);
117        context.set_prompt_cache_key(Some("prefix-abc".to_string()));
118        context.set_session_affinity_key(Some("conversation-abc".to_string()));
119
120        let responses = provider.stream_response(&context).collect::<Vec<_>>().await;
121        let captured = server.captured().await;
122
123        assert!(responses.iter().all(Result::is_ok), "{responses:?}");
124        assert!(responses.iter().any(|response| matches!(response, Ok(LlmResponse::Done { .. }))));
125        assert_eq!(captured.path, "/v1/chat/completions");
126        assert!(captured.body.get("prompt_cache_key").is_none());
127        assert!(captured.body.get("session_id").is_none());
128        assert!(captured.body.get("user").is_none());
129    }
130}