Skip to main content

starweaver_model/adapter/
traits.rs

1use async_trait::async_trait;
2use starweaver_usage::Usage;
3
4use crate::{
5    ModelResponse, message::ModelMessage, profile::ModelProfile, settings::ModelSettings,
6    stream::ModelResponseStreamEvent,
7};
8
9use super::{ModelError, ModelRequestContext, ModelRequestParameters, ModelResponseEventStream};
10
11/// Run-scoped model session used by agent loops that may issue multiple model requests.
12#[async_trait]
13pub trait ModelRunSession: Send {
14    /// Stream a model request and yield canonical events as they arrive.
15    async fn request_stream_incremental(
16        &mut self,
17        messages: Vec<ModelMessage>,
18        settings: Option<ModelSettings>,
19        params: ModelRequestParameters,
20        context: ModelRequestContext,
21    ) -> Result<ModelResponseEventStream, ModelError>;
22
23    /// Close run-scoped transport resources held by this session.
24    async fn close(&mut self) {}
25
26    /// Stream a model request and return the assembled final response.
27    async fn request_stream_final(
28        &mut self,
29        messages: Vec<ModelMessage>,
30        settings: Option<ModelSettings>,
31        params: ModelRequestParameters,
32        context: ModelRequestContext,
33    ) -> Result<ModelResponse, ModelError> {
34        let mut stream = self
35            .request_stream_incremental(messages, settings, params, context)
36            .await?;
37        while let Some(event) = stream.recv().await {
38            if let ModelResponseStreamEvent::FinalResult(response) = event? {
39                return Ok(*response);
40            }
41        }
42        Err(ModelError::UnsupportedResponse(
43            "model stream did not produce a final result".to_string(),
44        ))
45    }
46}
47
48struct DefaultModelRunSession<'a, M: ModelAdapter + ?Sized> {
49    model: &'a M,
50}
51
52#[async_trait]
53impl<M> ModelRunSession for DefaultModelRunSession<'_, M>
54where
55    M: ModelAdapter + ?Sized,
56{
57    async fn request_stream_incremental(
58        &mut self,
59        messages: Vec<ModelMessage>,
60        settings: Option<ModelSettings>,
61        params: ModelRequestParameters,
62        context: ModelRequestContext,
63    ) -> Result<ModelResponseEventStream, ModelError> {
64        self.model
65            .request_stream_incremental(messages, settings, params, context)
66            .await
67    }
68
69    async fn request_stream_final(
70        &mut self,
71        messages: Vec<ModelMessage>,
72        settings: Option<ModelSettings>,
73        params: ModelRequestParameters,
74        context: ModelRequestContext,
75    ) -> Result<ModelResponse, ModelError> {
76        self.model
77            .request_stream_final(messages, settings, params, context)
78            .await
79    }
80}
81
82/// Provider-neutral model adapter.
83#[async_trait]
84pub trait ModelAdapter: Send + Sync {
85    /// Provider model name.
86    fn model_name(&self) -> &str;
87
88    /// Provider name.
89    fn provider_name(&self) -> Option<&str>;
90
91    /// Model capability profile.
92    fn profile(&self) -> &ModelProfile;
93
94    /// Default generation settings.
95    fn default_settings(&self) -> Option<&ModelSettings>;
96
97    /// Create a run-scoped session for agent loops that may issue multiple model requests.
98    fn start_run_session(&self) -> Box<dyn ModelRunSession + '_> {
99        Box::new(DefaultModelRunSession { model: self })
100    }
101
102    /// Perform a complete model request.
103    async fn request(
104        &self,
105        messages: Vec<ModelMessage>,
106        settings: Option<ModelSettings>,
107        params: ModelRequestParameters,
108        context: ModelRequestContext,
109    ) -> Result<ModelResponse, ModelError>;
110
111    /// Stream a model request as canonical response part deltas.
112    async fn request_stream(
113        &self,
114        messages: Vec<ModelMessage>,
115        settings: Option<ModelSettings>,
116        params: ModelRequestParameters,
117        context: ModelRequestContext,
118    ) -> Result<Vec<ModelResponseStreamEvent>, ModelError> {
119        let response = self.request(messages, settings, params, context).await?;
120        Ok(vec![ModelResponseStreamEvent::FinalResult(Box::new(
121            response,
122        ))])
123    }
124
125    /// Stream a model request and return the assembled final response.
126    async fn request_stream_final(
127        &self,
128        messages: Vec<ModelMessage>,
129        settings: Option<ModelSettings>,
130        params: ModelRequestParameters,
131        context: ModelRequestContext,
132    ) -> Result<ModelResponse, ModelError> {
133        let events = self
134            .request_stream(messages, settings, params, context)
135            .await?;
136        events
137            .into_iter()
138            .find_map(|event| match event {
139                ModelResponseStreamEvent::FinalResult(response) => Some(*response),
140                ModelResponseStreamEvent::PartStart(_)
141                | ModelResponseStreamEvent::PartDelta(_)
142                | ModelResponseStreamEvent::PartEnd(_)
143                | ModelResponseStreamEvent::Diagnostic(_) => None,
144            })
145            .ok_or_else(|| {
146                ModelError::UnsupportedResponse(
147                    "model stream did not produce a final result".to_string(),
148                )
149            })
150    }
151
152    /// Stream a model request and yield canonical events as they arrive.
153    async fn request_stream_incremental(
154        &self,
155        messages: Vec<ModelMessage>,
156        settings: Option<ModelSettings>,
157        params: ModelRequestParameters,
158        context: ModelRequestContext,
159    ) -> Result<ModelResponseEventStream, ModelError> {
160        let cancellation_token = context.cancellation_token();
161        let events = self
162            .request_stream(messages, settings, params, context)
163            .await?;
164        let (sender, receiver) = tokio::sync::mpsc::channel(events.len().max(1));
165        for event in events {
166            let _ = sender.send(Ok(event)).await;
167        }
168        Ok(ModelResponseEventStream::new_with_cancellation(
169            receiver,
170            cancellation_token,
171        ))
172    }
173
174    /// Count tokens for a request where provider support exists.
175    async fn count_tokens(
176        &self,
177        _messages: &[ModelMessage],
178        _settings: Option<&ModelSettings>,
179        _params: &ModelRequestParameters,
180    ) -> Result<Usage, ModelError> {
181        Ok(Usage::default())
182    }
183}