starweaver_model/adapter/
traits.rs1use 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#[async_trait]
13pub trait ModelRunSession: Send {
14 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 async fn close(&mut self) {}
25
26 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#[async_trait]
84pub trait ModelAdapter: Send + Sync {
85 fn model_name(&self) -> &str;
87
88 fn provider_name(&self) -> Option<&str>;
90
91 fn profile(&self) -> &ModelProfile;
93
94 fn default_settings(&self) -> Option<&ModelSettings>;
96
97 fn start_run_session(&self) -> Box<dyn ModelRunSession + '_> {
99 Box::new(DefaultModelRunSession { model: self })
100 }
101
102 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 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 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 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 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}