返回 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 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
100 lines RUST