返回 CodeWhale
model.rs
根目录 / crates / tui / src / core / runtime_contract / model.rs
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 codewhale_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 pub trait ModelClient: Send + Sync {
18 fn provider_name(&self) -> &str;
19 fn model(&self) -> &str;
20 /// Concrete route base for billing classification, when this client can
21 /// prove one. Provider-neutral injected clients leave it unknown.
22 fn billing_base_url(&self) -> Option<&str> {
23 None
24 }
25 fn route_limits(&self) -> Option<codewhale_config::route::RouteLimits> {
26 None
27 }
28 fn effective_max_output_tokens(&self, requested_model: &str) -> u32 {
29 let route = self.effective_route_envelope(requested_model, chrono::Utc::now());
30 crate::route_budget::effective_max_output_tokens_for_route(
31 route.provider,
32 &route.model,
33 self.route_limits(),
34 )
35 }
36 fn effective_route_envelope(
37 &self,
38 requested_model: &str,
39 dispatched_at: chrono::DateTime<chrono::Utc>,
40 ) -> crate::cost_status::EffectiveRouteEnvelope {
41 let provider = crate::config::ProviderKind::parse(self.provider_name())
42 .unwrap_or(crate::config::ProviderKind::Custom);
43 crate::cost_status::EffectiveRouteEnvelope::capture_observed(
44 provider,
45 self.provider_name(),
46 requested_model,
47 self.billing_base_url(),
48 dispatched_at,
49 )
50 }
51 async fn create_message(&self, request: MessageRequest) -> Result<MessageResponse>;
52 /// Fresh authorization evidence; cache-owning adapters must bypass it.
53 async fn create_message_uncached(&self, request: MessageRequest) -> Result<MessageResponse> {
54 self.create_message(request).await
55 }
56 async fn create_message_stream(&self, request: MessageRequest) -> Result<StreamEventBox>;
57 #[expect(dead_code)]
58 async fn health_check(&self) -> Result<bool>;
59 }
60
61 pub type SharedModelClient = Arc<dyn ModelClient>;
62
63 /// Every existing provider client automatically satisfies the injectable
64 /// boundary. This keeps provider-specific HTTP/routing code behind
65 /// `LlmClient` while the Engine owns only the object-safe contract.
66 #[async_trait]
67 impl<T> ModelClient for T
68 where
69 T: LlmClient + Send + Sync,
70 {
71 fn provider_name(&self) -> &str {
72 LlmClient::provider_name(self)
73 }
74
75 fn model(&self) -> &str {
76 LlmClient::model(self)
77 }
78
79 fn billing_base_url(&self) -> Option<&str> {
80 LlmClient::billing_base_url(self)
81 }
82
83 fn route_limits(&self) -> Option<codewhale_config::route::RouteLimits> {
84 LlmClient::route_limits(self)
85 }
86
87 fn effective_max_output_tokens(&self, requested_model: &str) -> u32 {
88 LlmClient::effective_max_output_tokens(self, requested_model)
89 }
90
91 fn effective_route_envelope(
92 &self,
93 requested_model: &str,
94 dispatched_at: chrono::DateTime<chrono::Utc>,
95 ) -> crate::cost_status::EffectiveRouteEnvelope {
96 LlmClient::effective_route_envelope(self, requested_model, dispatched_at)
97 }
98
99 async fn create_message(&self, request: MessageRequest) -> Result<MessageResponse> {
100 LlmClient::create_message(self, request).await
101 }
102
103 async fn create_message_uncached(&self, request: MessageRequest) -> Result<MessageResponse> {
104 LlmClient::create_message_uncached(self, request).await
105 }
106
107 async fn create_message_stream(&self, request: MessageRequest) -> Result<StreamEventBox> {
108 LlmClient::create_message_stream(self, request).await
109 }
110
111 async fn health_check(&self) -> Result<bool> {
112 LlmClient::health_check(self).await
113 }
114 }
115
116 #[cfg(test)]
117 mod tests {
118 use super::*;
119
120 #[test]
121 fn model_client_is_object_safe() {
122 fn accepts_dyn(_: Option<SharedModelClient>) {}
123 accepts_dyn(None);
124 }
125 }
126
126 lines RUST