1use std::fmt;
24
25use crate::error::LlmError;
26use crate::openai::{CompletionTokensParam, OpenAiConfig, OpenAiProvider};
27use crate::provider::{
28 ChatExtras, ChatResponse, ChatStream, GenerationOverrides, LlmProvider, Message, StatusTx,
29 ToolDefinition,
30};
31
32#[derive(Clone)]
55pub struct CompatibleConfig {
56 pub provider_name: String,
58 pub api_key: String,
60 pub base_url: String,
62 pub model: String,
64 pub max_tokens: u32,
66 pub embedding_model: Option<String>,
68 pub completion_tokens_param: Option<CompletionTokensParam>,
74 pub vision: Option<bool>,
81}
82
83impl fmt::Debug for CompatibleConfig {
84 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
85 f.debug_struct("CompatibleConfig")
86 .field("provider_name", &self.provider_name)
87 .field("api_key", &"<redacted>")
88 .field("base_url", &self.base_url)
89 .field("model", &self.model)
90 .field("max_tokens", &self.max_tokens)
91 .field("embedding_model", &self.embedding_model)
92 .field("completion_tokens_param", &self.completion_tokens_param)
93 .field("vision", &self.vision)
94 .finish()
95 }
96}
97
98pub struct CompatibleProvider {
103 inner: OpenAiProvider,
104 provider_name: String,
106}
107
108impl CompatibleProvider {
109 #[must_use]
111 pub fn new(cfg: CompatibleConfig) -> Self {
112 let provider_name = cfg.provider_name;
113 let inner = OpenAiProvider::new(OpenAiConfig {
114 api_key: cfg.api_key,
115 base_url: cfg.base_url,
116 model: cfg.model,
117 max_tokens: cfg.max_tokens,
118 embedding_model: cfg.embedding_model,
119 reasoning_effort: None,
120 context_window: None,
121 completion_tokens_param: cfg.completion_tokens_param,
122 vision: cfg.vision,
123 });
124 Self {
125 inner,
126 provider_name,
127 }
128 }
129}
130
131impl fmt::Debug for CompatibleProvider {
132 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
133 f.debug_struct("CompatibleProvider")
134 .field("provider_name", &self.provider_name)
135 .field("inner", &self.inner)
136 .finish_non_exhaustive()
137 }
138}
139
140impl Clone for CompatibleProvider {
141 fn clone(&self) -> Self {
142 Self {
143 inner: self.inner.clone(),
144 provider_name: self.provider_name.clone(),
145 }
146 }
147}
148
149impl CompatibleProvider {
150 pub async fn list_models_remote(
156 &self,
157 ) -> Result<Vec<crate::model_cache::RemoteModelInfo>, LlmError> {
158 self.inner.list_models_remote().await
159 }
160}
161
162impl CompatibleProvider {
163 pub fn set_status_tx(&mut self, tx: StatusTx) {
165 self.inner.status_tx = Some(tx);
166 }
167
168 #[must_use]
170 pub fn with_generation_overrides(mut self, overrides: GenerationOverrides) -> Self {
171 self.inner = self.inner.with_generation_overrides(overrides);
172 self
173 }
174
175 #[must_use]
199 pub fn with_completion_tokens_param(mut self, param: CompletionTokensParam) -> Self {
200 self.inner = self.inner.with_completion_tokens_param(param);
201 self
202 }
203
204 #[must_use]
231 pub fn with_vision(mut self, supported: bool) -> Self {
232 self.inner = self.inner.with_vision(supported);
233 self
234 }
235
236 #[must_use]
242 pub fn with_output_schema_forwarding(
243 mut self,
244 enabled: bool,
245 hint_bytes: usize,
246 max_description_bytes: usize,
247 ) -> Self {
248 self.inner =
249 self.inner
250 .with_output_schema_forwarding(enabled, hint_bytes, max_description_bytes);
251 self
252 }
253
254 pub fn set_reasoning_effort(&mut self, effort: Option<String>) {
260 self.inner.set_reasoning_effort(effort);
261 }
262
263 #[must_use]
266 pub fn current_reasoning_effort(&self) -> Option<String> {
267 self.inner.reasoning_effort.clone()
268 }
269}
270
271impl LlmProvider for CompatibleProvider {
272 fn context_window(&self) -> Option<usize> {
273 self.inner.context_window()
274 }
275
276 #[tracing::instrument(
277 name = "llm.chat",
278 skip_all,
279 fields(provider = self.name(), model = self.model_identifier())
280 )]
281 async fn chat(&self, messages: &[Message]) -> Result<String, LlmError> {
282 self.inner.chat(messages).await
283 }
284
285 async fn chat_with_extras(
286 &self,
287 messages: &[Message],
288 ) -> Result<(String, ChatExtras), LlmError> {
289 self.inner.chat_with_extras(messages).await
290 }
291
292 #[tracing::instrument(
293 name = "llm.chat_stream",
294 skip_all,
295 fields(provider = self.name(), model = self.model_identifier())
296 )]
297 async fn chat_stream(&self, messages: &[Message]) -> Result<ChatStream, LlmError> {
298 self.inner.chat_stream(messages).await
299 }
300
301 fn supports_streaming(&self) -> bool {
302 self.inner.supports_streaming()
303 }
304
305 #[tracing::instrument(
306 name = "llm.embed",
307 skip_all,
308 fields(provider = self.name(), model = self.model_identifier())
309 )]
310 async fn embed(&self, text: &str) -> Result<Vec<f32>, LlmError> {
311 self.inner.embed(text).await
312 }
313
314 #[tracing::instrument(
315 name = "llm.embed_batch",
316 skip_all,
317 fields(provider = self.name(), model = self.model_identifier())
318 )]
319 async fn embed_batch(&self, texts: &[&str]) -> Result<Vec<Vec<f32>>, LlmError> {
320 self.inner.embed_batch(texts).await
321 }
322
323 fn supports_embeddings(&self) -> bool {
324 self.inner.supports_embeddings()
325 }
326
327 fn name(&self) -> &str {
328 &self.provider_name
329 }
330
331 fn model_identifier(&self) -> &str {
332 self.inner.model_identifier()
333 }
334
335 fn list_models(&self) -> Vec<String> {
336 self.inner.list_models()
337 }
338
339 fn supports_structured_output(&self) -> bool {
340 self.inner.supports_structured_output()
341 }
342
343 async fn chat_typed<T>(&self, messages: &[Message]) -> Result<T, LlmError>
344 where
345 T: serde::de::DeserializeOwned + schemars::JsonSchema + 'static,
346 Self: Sized,
347 {
348 self.inner.chat_typed(messages).await
349 }
350
351 #[tracing::instrument(
352 name = "llm.chat_with_tools",
353 skip_all,
354 fields(provider = self.name(), model = self.model_identifier(), tool_count = tools.len())
355 )]
356 async fn chat_with_tools(
357 &self,
358 messages: &[Message],
359 tools: &[ToolDefinition],
360 ) -> Result<ChatResponse, LlmError> {
361 self.inner.chat_with_tools(messages, tools).await
362 }
363
364 fn last_cache_usage(&self) -> Option<(u64, u64)> {
365 self.inner.last_cache_usage()
366 }
367
368 fn last_usage(&self) -> Option<(u64, u64)> {
369 self.inner.last_usage()
370 }
371
372 fn last_reasoning_tokens(&self) -> Option<u64> {
373 self.inner.last_reasoning_tokens()
374 }
375
376 fn supports_vision(&self) -> bool {
377 self.inner.supports_vision()
378 }
379
380 fn supports_tool_use(&self) -> bool {
381 self.inner.supports_tool_use()
382 }
383
384 fn debug_request_json(
385 &self,
386 messages: &[Message],
387 tools: &[ToolDefinition],
388 stream: bool,
389 ) -> serde_json::Value {
390 self.inner.debug_request_json(messages, tools, stream)
391 }
392}
393
394#[cfg(test)]
395mod tests {
396 use super::*;
397
398 fn test_provider() -> CompatibleProvider {
399 CompatibleProvider::new(CompatibleConfig {
400 provider_name: "groq".into(),
401 api_key: "key".into(),
402 base_url: "https://api.groq.com/openai/v1".into(),
403 model: "llama-3.3-70b".into(),
404 max_tokens: 4096,
405 embedding_model: None,
406 completion_tokens_param: None,
407 vision: None,
408 })
409 }
410
411 #[test]
412 fn name_returns_custom_provider_name() {
413 let p = test_provider();
414 assert_eq!(p.name(), "groq");
415 }
416
417 #[test]
418 fn context_window_delegates_to_inner() {
419 let p = CompatibleProvider::new(CompatibleConfig {
421 provider_name: "openai".into(),
422 api_key: "key".into(),
423 base_url: "https://api.openai.com/v1".into(),
424 model: "gpt-4o".into(),
425 max_tokens: 4096,
426 embedding_model: None,
427 completion_tokens_param: None,
428 vision: None,
429 });
430 assert_eq!(p.context_window(), Some(128_000));
431 }
432
433 #[test]
434 fn context_window_unknown_model_returns_some_fallback() {
435 let p = CompatibleProvider::new(CompatibleConfig {
437 provider_name: "local".into(),
438 api_key: "key".into(),
439 base_url: "http://localhost/v1".into(),
440 model: "unknown-custom-model".into(),
441 max_tokens: 4096,
442 embedding_model: None,
443 completion_tokens_param: None,
444 vision: None,
445 });
446 assert!(p.context_window().is_some());
448 }
449
450 #[test]
451 fn supports_streaming_delegates() {
452 assert!(test_provider().supports_streaming());
453 }
454
455 #[test]
456 fn supports_embeddings_without_model() {
457 assert!(!test_provider().supports_embeddings());
458 }
459
460 #[test]
461 fn supports_embeddings_with_model() {
462 let p = CompatibleProvider::new(CompatibleConfig {
463 provider_name: "test".into(),
464 api_key: "key".into(),
465 base_url: "http://localhost".into(),
466 model: "m".into(),
467 max_tokens: 100,
468 embedding_model: Some("embed-model".into()),
469 completion_tokens_param: None,
470 vision: None,
471 });
472 assert!(p.supports_embeddings());
473 }
474
475 #[test]
476 fn clone_preserves_name() {
477 let p = test_provider();
478 let c = p.clone();
479 assert_eq!(c.name(), "groq");
480 }
481
482 #[test]
483 fn debug_contains_provider_name() {
484 let debug = format!("{:?}", test_provider());
485 assert!(debug.contains("groq"));
486 assert!(debug.contains("CompatibleProvider"));
487 }
488
489 #[tokio::test]
490 async fn chat_unreachable_errors() {
491 let p = CompatibleProvider::new(CompatibleConfig {
492 provider_name: "test".into(),
493 api_key: "key".into(),
494 base_url: "http://127.0.0.1:1".into(),
495 model: "m".into(),
496 max_tokens: 100,
497 embedding_model: None,
498 completion_tokens_param: None,
499 vision: None,
500 });
501 let msgs = vec![Message::from_legacy(crate::provider::Role::User, "hello")];
502 assert!(p.chat(&msgs).await.is_err());
503 }
504
505 #[tokio::test]
506 async fn embed_without_model_errors() {
507 let p = test_provider();
508 let result = p.embed("test").await;
509 assert!(result.is_err());
510 }
511
512 #[test]
513 fn last_usage_initially_none() {
514 assert!(test_provider().last_usage().is_none());
515 }
516
517 #[test]
518 fn with_output_schema_forwarding_does_not_panic() {
519 let p = test_provider().with_output_schema_forwarding(true, 512, usize::MAX);
521 assert_eq!(p.name(), "groq");
522 }
523
524 #[test]
527 fn set_reasoning_effort_applies_via_compatible() {
528 let mut p = test_provider();
529 p.set_reasoning_effort(Some("high".into()));
530 assert_eq!(p.inner.reasoning_effort.as_deref(), Some("high"));
531 }
532
533 #[test]
534 fn any_provider_set_reasoning_effort_delegates_to_compatible() {
535 use crate::any::AnyProvider;
536 let mut any = AnyProvider::Compatible(test_provider());
537 any.set_reasoning_effort(Some("high".into()));
538 let AnyProvider::Compatible(ref p) = any else {
539 panic!("variant must remain Compatible");
540 };
541 assert_eq!(
542 p.inner.reasoning_effort.as_deref(),
543 Some("high"),
544 "Compatible inner OpenAiProvider must have reasoning_effort applied"
545 );
546 }
547
548 #[test]
549 fn supports_vision_delegates_to_inner() {
550 assert!(!test_provider().supports_vision());
553 }
554
555 #[test]
556 fn supports_vision_with_vision_override_delegates_to_inner() {
557 let p = test_provider().with_vision(true);
560 assert!(p.supports_vision());
561 }
562
563 #[test]
564 fn supports_vision_config_field_true_forwards_to_inner() {
565 let p = CompatibleProvider::new(CompatibleConfig {
568 provider_name: "test".into(),
569 api_key: "key".into(),
570 base_url: "http://localhost".into(),
571 model: "llama-3.3-70b".into(),
572 max_tokens: 100,
573 embedding_model: None,
574 completion_tokens_param: None,
575 vision: Some(true),
576 });
577 assert!(p.supports_vision());
578 }
579
580 #[test]
581 fn supports_vision_config_field_false_forwards_to_inner() {
582 let p = CompatibleProvider::new(CompatibleConfig {
583 provider_name: "test".into(),
584 api_key: "key".into(),
585 base_url: "http://localhost".into(),
586 model: "llama-3.3-70b".into(),
587 max_tokens: 100,
588 embedding_model: None,
589 completion_tokens_param: None,
590 vision: Some(false),
591 });
592 assert!(!p.supports_vision());
593 }
594
595 #[test]
596 fn supports_tool_use_delegates_to_inner() {
597 assert!(test_provider().supports_tool_use());
599 }
600
601 #[test]
602 fn last_reasoning_tokens_initially_none() {
603 assert!(test_provider().last_reasoning_tokens().is_none());
604 }
605
606 #[test]
607 fn compatible_config_debug_redacts_api_key() {
608 let cfg = CompatibleConfig {
609 provider_name: "together-ai".into(),
610 api_key: "sk-SUPERSECRET".into(),
611 base_url: "https://api.together.xyz/v1".into(),
612 model: "meta-llama/Llama-3.3-70B-Instruct-Turbo".into(),
613 max_tokens: 4096,
614 embedding_model: None,
615 completion_tokens_param: None,
616 vision: None,
617 };
618 let dbg = format!("{cfg:?}");
619 assert!(!dbg.contains("sk-SUPERSECRET"));
620 assert!(dbg.contains("<redacted>"));
621 }
622}