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::error_stream;
6use crate::providers::openai_compatible::create_custom_stream_generic;
7use crate::{Context, LlmResponseStream, StreamingModelProvider};
8
9/// A Provider that's compatible with `OpenAI`'s chat completion API
10/// Other providers (e.g. Ollama, Llama.cpp etc) that are "`OpenAI` compatible" should implement this trait
11pub trait OpenAiChatProvider {
12    type Config: Config + Clone + 'static;
13
14    fn client(&self) -> &Client<Self::Config>;
15    fn model(&self) -> &str;
16    fn provider_name(&self) -> &str;
17}
18
19impl<T: OpenAiChatProvider + Send + Sync> StreamingModelProvider for T {
20    fn stream_response(&self, context: &Context) -> LlmResponseStream {
21        if let Err(error) = crate::provider::validate_reasoning(context, None) {
22            return crate::provider::error_stream(error);
23        }
24        let model = self.model().to_string();
25        let messages = match map_messages(context.messages()) {
26            Ok(messages) => messages,
27            Err(e) => return error_stream(e),
28        };
29        let message_count = messages.len();
30        let tools = if context.tools().is_empty() {
31            None
32        } else {
33            match map_tools(context.tools(), None) {
34                Ok(t) => Some(t),
35                Err(e) => return error_stream(e),
36            }
37        };
38
39        debug!("Starting chat completion stream for model: {model} with {message_count} messages");
40        let request = CreateChatCompletionRequest { model, messages, tools, stream: Some(true), ..Default::default() };
41        create_custom_stream_generic(self.client(), request)
42    }
43
44    fn context_window(&self) -> Option<u32> {
45        None
46    }
47
48    fn display_name(&self) -> String {
49        let model = self.model();
50        if model.is_empty() { self.provider_name().to_string() } else { format!("{} ({model})", self.provider_name()) }
51    }
52}
53
54#[cfg(test)]
55mod tests {
56    use futures::StreamExt;
57
58    use super::*;
59    use crate::providers::local::ollama::OllamaProvider;
60    use crate::providers::test_capture_server::CaptureServer;
61    use crate::{ChatMessage, LlmResponse, Result};
62
63    #[tokio::test]
64    async fn local_chat_providers_omit_cache_metadata() {
65        let mut server = CaptureServer::start_chat_completions().await;
66        let provider = OllamaProvider::new("test-model", &server.base_url);
67        let mut context = Context::new(vec![ChatMessage::user("Hello")], vec![]);
68        context.set_prompt_cache_key(Some("prefix-abc".to_string()));
69        context.set_session_affinity_key(Some("conversation-abc".to_string()));
70
71        let responses = provider.stream_response(&context).collect::<Vec<_>>().await;
72        let captured = server.captured().await;
73
74        assert!(responses.iter().all(Result::is_ok), "{responses:?}");
75        assert!(responses.iter().any(|response| matches!(response, Ok(LlmResponse::Done { .. }))));
76        assert_eq!(captured.path, "/v1/chat/completions");
77        assert!(captured.body.get("prompt_cache_key").is_none());
78        assert!(captured.body.get("session_id").is_none());
79        assert!(captured.body.get("user").is_none());
80    }
81}