Skip to main content

llm/providers/openai/
provider.rs

1use async_openai::{Client, config::Config, types::chat::CreateChatCompletionRequest};
2use tracing::debug;
3
4use super::mappers::{map_messages, map_tools};
5use crate::provider::validate_reasoning;
6use crate::providers::openai_compatible::create_custom_stream_generic;
7use crate::providers::response_stream::error_stream;
8use crate::{Context, LlmResponseStream, Result, StreamingModelProvider};
9use std::time::Duration;
10
11/// A Provider that's compatible with `OpenAI`'s chat completion API
12/// Other providers (e.g. Ollama, Llama.cpp etc) that are "`OpenAI` compatible" should implement this trait
13pub trait OpenAiChatProvider {
14    type Config: Config + Clone + 'static;
15
16    fn client(&self) -> &Client<Self::Config>;
17    fn model(&self) -> &str;
18    fn provider_name(&self) -> &str;
19    fn idle_timeout(&self) -> Duration;
20}
21
22impl<T: OpenAiChatProvider + Send + Sync> StreamingModelProvider for T {
23    fn stream_response(&self, context: &Context) -> LlmResponseStream {
24        try_stream_response(self, context).unwrap_or_else(error_stream)
25    }
26
27    fn context_window(&self) -> Option<u32> {
28        None
29    }
30
31    fn display_name(&self) -> String {
32        let model = self.model();
33        if model.is_empty() { self.provider_name().to_string() } else { format!("{} ({model})", self.provider_name()) }
34    }
35}
36
37fn try_stream_response<T: OpenAiChatProvider>(provider: &T, context: &Context) -> Result<LlmResponseStream> {
38    validate_reasoning(context, None)?;
39    let model = provider.model().to_string();
40    let messages = map_messages(context.messages())?;
41    let message_count = messages.len();
42    let tools = if context.tools().is_empty() { None } else { Some(map_tools(context.tools(), None)?) };
43
44    debug!("Starting chat completion stream for model: {model} with {message_count} messages");
45    let request = CreateChatCompletionRequest { model, messages, tools, stream: Some(true), ..Default::default() };
46    Ok(create_custom_stream_generic(provider.client(), request, provider.idle_timeout()))
47}
48
49#[cfg(test)]
50mod tests {
51    use futures::StreamExt;
52
53    use super::*;
54    use crate::providers::local::ollama::OllamaProvider;
55    use crate::providers::test_capture_server::{CaptureServer, ResponseSpec, hello_context};
56    use crate::{LlmError, LlmResponse, ProviderConnectionConfig, ProviderErrorKind, ProviderFactory};
57
58    #[tokio::test]
59    async fn local_provider_honors_configured_idle_timeout() {
60        let spec = ResponseSpec::sse(include_str!("../../../tests/fixtures/openai/01_minimal.sse"))
61            .paced(Duration::from_mins(2));
62        let mut server = CaptureServer::start_with_spec(spec).await;
63        let connection = ProviderConnectionConfig {
64            base_url: Some(server.base_url.clone()),
65            idle_timeout: Duration::from_mins(1),
66            ..Default::default()
67        };
68        let provider = OllamaProvider::from_env_with_connection(connection).await.unwrap().with_model("test-model");
69
70        let responses = server.collect_on_paused_clock(provider.stream_response(&hello_context())).await;
71
72        let error = responses.last().and_then(|response| response.as_ref().err()).and_then(LlmError::provider);
73        assert_eq!(error.map(|error| error.kind), Some(ProviderErrorKind::Timeout), "{responses:?}");
74    }
75
76    #[tokio::test]
77    async fn local_chat_providers_omit_cache_metadata() {
78        let mut server = CaptureServer::start_chat_completions().await;
79        let provider = OllamaProvider::new("test-model", &server.base_url);
80        let mut context = hello_context();
81        context.set_prompt_cache_key(Some("prefix-abc".to_string()));
82        context.set_session_affinity_key(Some("conversation-abc".to_string()));
83
84        let responses = provider.stream_response(&context).collect::<Vec<_>>().await;
85        let captured = server.captured().await;
86
87        assert!(responses.iter().all(Result::is_ok), "{responses:?}");
88        assert!(responses.iter().any(|response| matches!(response, Ok(LlmResponse::Done { .. }))));
89        assert_eq!(captured.path, "/v1/chat/completions");
90        assert!(captured.body.get("prompt_cache_key").is_none());
91        assert!(captured.body.get("session_id").is_none());
92        assert!(captured.body.get("user").is_none());
93    }
94}