| 1 | use std::sync::Arc; |
| 2 | |
| 3 | use anyhow::Result; |
| 4 | use async_trait::async_trait; |
| 5 | |
| 6 | use crate::llm_client::LlmClient; |
| 7 | use crate::llm_client::StreamEventBox; |
| 8 | use crate::models::{MessageRequest, MessageResponse}; |
| 9 | |
| 10 | /// Object-safe model boundary for Engine dependency injection. |
| 11 | /// |
| 12 | /// The existing `LlmClient` uses return-position `impl Future`, which is |
| 13 | /// efficient for concrete providers but cannot be placed behind `dyn`. This |
| 14 | /// adapter preserves that provider trait while giving deterministic Engine |
| 15 | /// tests and alternate adapters one injectable boundary. |
| 16 | #[async_trait] |
| 17 | #[allow(dead_code)] |
| 18 | pub trait ModelClient: Send + Sync { |
| 19 | fn provider_name(&self) -> &str; |
| 20 | fn model(&self) -> &str; |
| 21 | /// Concrete route base for billing classification, when this client can |
| 22 | /// prove one. Provider-neutral injected clients leave it unknown. |
| 23 | fn billing_base_url(&self) -> Option<&str> { |
| 24 | None |
| 25 | } |
| 26 | fn effective_route_envelope( |
| 27 | &self, |
| 28 | requested_model: &str, |
| 29 | dispatched_at: chrono::DateTime<chrono::Utc>, |
| 30 | ) -> crate::cost_status::EffectiveRouteEnvelope { |
| 31 | let provider = crate::config::ApiProvider::parse(self.provider_name()) |
| 32 | .unwrap_or(crate::config::ApiProvider::Custom); |
| 33 | crate::cost_status::EffectiveRouteEnvelope::capture( |
| 34 | None, |
| 35 | provider, |
| 36 | self.provider_name(), |
| 37 | requested_model, |
| 38 | self.billing_base_url(), |
| 39 | dispatched_at, |
| 40 | ) |
| 41 | } |
| 42 | async fn create_message(&self, request: MessageRequest) -> Result<MessageResponse>; |
| 43 | async fn create_message_stream(&self, request: MessageRequest) -> Result<StreamEventBox>; |
| 44 | async fn health_check(&self) -> Result<bool>; |
| 45 | } |
| 46 | |
| 47 | pub type SharedModelClient = Arc<dyn ModelClient>; |
| 48 | |
| 49 | /// Every existing provider client automatically satisfies the injectable |
| 50 | /// boundary. This keeps provider-specific HTTP/routing code behind |
| 51 | /// `LlmClient` while the Engine owns only the object-safe contract. |
| 52 | #[async_trait] |
| 53 | impl<T> ModelClient for T |
| 54 | where |
| 55 | T: LlmClient + Send + Sync, |
| 56 | { |
| 57 | fn provider_name(&self) -> &str { |
| 58 | LlmClient::provider_name(self) |
| 59 | } |
| 60 | |
| 61 | fn model(&self) -> &str { |
| 62 | LlmClient::model(self) |
| 63 | } |
| 64 | |
| 65 | fn billing_base_url(&self) -> Option<&str> { |
| 66 | LlmClient::billing_base_url(self) |
| 67 | } |
| 68 | |
| 69 | fn effective_route_envelope( |
| 70 | &self, |
| 71 | requested_model: &str, |
| 72 | dispatched_at: chrono::DateTime<chrono::Utc>, |
| 73 | ) -> crate::cost_status::EffectiveRouteEnvelope { |
| 74 | LlmClient::effective_route_envelope(self, requested_model, dispatched_at) |
| 75 | } |
| 76 | |
| 77 | async fn create_message(&self, request: MessageRequest) -> Result<MessageResponse> { |
| 78 | LlmClient::create_message(self, request).await |
| 79 | } |
| 80 | |
| 81 | async fn create_message_stream(&self, request: MessageRequest) -> Result<StreamEventBox> { |
| 82 | LlmClient::create_message_stream(self, request).await |
| 83 | } |
| 84 | |
| 85 | async fn health_check(&self) -> Result<bool> { |
| 86 | LlmClient::health_check(self).await |
| 87 | } |
| 88 | } |
| 89 | |
| 90 | #[cfg(test)] |
| 91 | mod tests { |
| 92 | use super::*; |
| 93 | |
| 94 | #[test] |
| 95 | fn model_client_is_object_safe() { |
| 96 | fn accepts_dyn(_: Option<SharedModelClient>) {} |
| 97 | accepts_dyn(None); |
| 98 | } |
| 99 | } |
| 100 |