telos_agent/model/provider/
traits.rs1use async_stream::try_stream;
4use async_trait::async_trait;
5use futures_core::stream::Stream;
6
7use crate::error::AgentError;
8
9use super::types::{CompletionRequest, CompletionResponse, ProviderEvent};
10
11#[async_trait]
13pub trait ModelProvider: Send + Sync {
14 async fn complete(&self, request: CompletionRequest) -> Result<CompletionResponse, AgentError>;
16
17 fn stream_complete<'a>(
23 &'a self,
24 request: CompletionRequest,
25 ) -> std::pin::Pin<Box<dyn Stream<Item = Result<ProviderEvent, AgentError>> + Send + 'a>> {
26 Box::pin(try_stream! {
27 let response = self.complete(request).await?;
28 yield ProviderEvent::MessageStart;
29 for block in &response.message.blocks {
30 match block {
31 crate::model::message::ContentBlock::Text(text) => {
32 yield ProviderEvent::TextDelta(text.text.clone());
33 }
34 crate::model::message::ContentBlock::Thinking(thinking) => {
35 yield ProviderEvent::ThinkingDelta(thinking.text.clone());
36 }
37 crate::model::message::ContentBlock::ToolCall(call) => {
38 yield ProviderEvent::ToolCall(call.clone());
39 }
40 crate::model::message::ContentBlock::ToolResult(_) => {}
41 }
42 }
43 yield ProviderEvent::MessageStop {
44 stop_reason: response.stop_reason,
45 usage: response.usage,
46 model: response.model,
47 };
48 })
49 }
50
51 fn max_tokens(&self) -> u32 {
57 128_000
58 }
59
60 fn estimate_tokens(&self, text: &str) -> usize {
70 crate::model::tokens::count_tokens(text)
71 }
72}
73
74pub struct ErasedProvider<'a>(pub &'a (dyn ModelProvider + 'a));
78
79#[async_trait]
80impl ModelProvider for ErasedProvider<'_> {
81 async fn complete(&self, request: CompletionRequest) -> Result<CompletionResponse, AgentError> {
82 self.0.complete(request).await
83 }
84
85 fn stream_complete<'a>(
86 &'a self,
87 request: CompletionRequest,
88 ) -> std::pin::Pin<Box<dyn Stream<Item = Result<ProviderEvent, AgentError>> + Send + 'a>> {
89 self.0.stream_complete(request)
90 }
91
92 fn max_tokens(&self) -> u32 {
93 self.0.max_tokens()
94 }
95
96 fn estimate_tokens(&self, text: &str) -> usize {
97 self.0.estimate_tokens(text)
98 }
99}