starweaver_model/wrappers/
concurrency.rs1use std::sync::Arc;
4
5use async_trait::async_trait;
6use serde_json::json;
7use tokio::sync::{OwnedSemaphorePermit, Semaphore};
8
9use super::DynModelAdapter;
10use crate::{
11 adapter::{
12 ModelAdapter, ModelError, ModelRequestContext, ModelRequestParameters,
13 ModelResponseEventStream, ModelRunSession,
14 },
15 message::{ModelMessage, ModelResponse},
16 profile::ModelProfile,
17 settings::ModelSettings,
18 stream::ModelResponseStreamEvent,
19};
20
21pub struct ConcurrencyLimitedModel {
23 inner: DynModelAdapter,
24 semaphore: Arc<Semaphore>,
25 max_concurrency: usize,
26}
27
28struct ConcurrencyLimitedRunSession<'a> {
29 model: &'a ConcurrencyLimitedModel,
30 inner: Box<dyn ModelRunSession + 'a>,
31}
32
33impl ConcurrencyLimitedModel {
34 #[must_use]
36 pub fn new(inner: DynModelAdapter, max_concurrency: usize) -> Self {
37 let permits = max_concurrency.max(1);
38 Self {
39 inner,
40 semaphore: Arc::new(Semaphore::new(permits)),
41 max_concurrency: permits,
42 }
43 }
44
45 #[must_use]
47 pub fn with_shared_semaphore(inner: DynModelAdapter, semaphore: Arc<Semaphore>) -> Self {
48 Self {
49 inner,
50 max_concurrency: semaphore.available_permits().max(1),
51 semaphore,
52 }
53 }
54
55 #[must_use]
57 pub const fn max_concurrency(&self) -> usize {
58 self.max_concurrency
59 }
60
61 async fn acquire(&self) -> Result<OwnedSemaphorePermit, ModelError> {
62 self.semaphore
63 .clone()
64 .acquire_owned()
65 .await
66 .map_err(|error| {
67 ModelError::Transport(format!("model concurrency limiter closed: {error}"))
68 })
69 }
70}
71
72#[async_trait]
73impl ModelAdapter for ConcurrencyLimitedModel {
74 fn model_name(&self) -> &str {
75 self.inner.model_name()
76 }
77
78 fn provider_name(&self) -> Option<&str> {
79 self.inner.provider_name()
80 }
81
82 fn profile(&self) -> &ModelProfile {
83 self.inner.profile()
84 }
85
86 fn default_settings(&self) -> Option<&ModelSettings> {
87 self.inner.default_settings()
88 }
89
90 fn start_run_session(&self) -> Box<dyn ModelRunSession + '_> {
91 Box::new(ConcurrencyLimitedRunSession {
92 model: self,
93 inner: self.inner.start_run_session(),
94 })
95 }
96
97 async fn request(
98 &self,
99 messages: Vec<ModelMessage>,
100 settings: Option<ModelSettings>,
101 params: ModelRequestParameters,
102 mut context: ModelRequestContext,
103 ) -> Result<ModelResponse, ModelError> {
104 let _permit = self.acquire().await?;
105 annotate_limiter(&mut context, self.max_concurrency);
106 self.inner
107 .request(messages, settings, params, context)
108 .await
109 }
110
111 async fn request_stream(
112 &self,
113 messages: Vec<ModelMessage>,
114 settings: Option<ModelSettings>,
115 params: ModelRequestParameters,
116 mut context: ModelRequestContext,
117 ) -> Result<Vec<ModelResponseStreamEvent>, ModelError> {
118 let _permit = self.acquire().await?;
119 annotate_limiter(&mut context, self.max_concurrency);
120 self.inner
121 .request_stream(messages, settings, params, context)
122 .await
123 }
124
125 async fn request_stream_incremental(
126 &self,
127 messages: Vec<ModelMessage>,
128 settings: Option<ModelSettings>,
129 params: ModelRequestParameters,
130 mut context: ModelRequestContext,
131 ) -> Result<ModelResponseEventStream, ModelError> {
132 let permit = self.acquire().await?;
133 annotate_limiter(&mut context, self.max_concurrency);
134 let events = self
135 .inner
136 .request_stream(messages, settings, params, context)
137 .await?;
138 let (sender, receiver) = tokio::sync::mpsc::channel(events.len().max(1));
139 tokio::spawn(async move {
140 let _permit = permit;
141 for event in events {
142 if sender.send(Ok(event)).await.is_err() {
143 return;
144 }
145 }
146 });
147 Ok(ModelResponseEventStream::new(receiver))
148 }
149
150 async fn count_tokens(
151 &self,
152 messages: &[ModelMessage],
153 settings: Option<&ModelSettings>,
154 params: &ModelRequestParameters,
155 ) -> Result<starweaver_usage::Usage, ModelError> {
156 self.inner.count_tokens(messages, settings, params).await
157 }
158}
159
160#[async_trait]
161impl ModelRunSession for ConcurrencyLimitedRunSession<'_> {
162 async fn request_stream_incremental(
163 &mut self,
164 messages: Vec<ModelMessage>,
165 settings: Option<ModelSettings>,
166 params: ModelRequestParameters,
167 mut context: ModelRequestContext,
168 ) -> Result<ModelResponseEventStream, ModelError> {
169 let cancellation_token = context.cancellation_token();
170 let permit = self.model.acquire().await?;
171 annotate_limiter(&mut context, self.model.max_concurrency);
172 let mut events = self
173 .inner
174 .request_stream_incremental(messages, settings, params, context)
175 .await?;
176 let drop_abort_token = events.drop_abort_token();
177 let (sender, receiver) = tokio::sync::mpsc::channel(32);
178 tokio::spawn(async move {
179 let _permit = permit;
180 while let Some(event) = events.recv().await {
181 if sender.send(event).await.is_err() {
182 return;
183 }
184 }
185 });
186 Ok(
187 ModelResponseEventStream::new_with_cancellation_and_drop_abort(
188 receiver,
189 cancellation_token,
190 drop_abort_token,
191 ),
192 )
193 }
194
195 async fn request_stream_final(
196 &mut self,
197 messages: Vec<ModelMessage>,
198 settings: Option<ModelSettings>,
199 params: ModelRequestParameters,
200 mut context: ModelRequestContext,
201 ) -> Result<ModelResponse, ModelError> {
202 let _permit = self.model.acquire().await?;
203 annotate_limiter(&mut context, self.model.max_concurrency);
204 self.inner
205 .request_stream_final(messages, settings, params, context)
206 .await
207 }
208
209 async fn close(&mut self) {
210 self.inner.close().await;
211 }
212}
213
214fn annotate_limiter(context: &mut ModelRequestContext, max_concurrency: usize) {
215 context.llm_trace_metadata.insert(
216 "starweaver_model_wrapper".to_string(),
217 json!({
218 "kind": "concurrency_limited",
219 "max_concurrency": max_concurrency,
220 }),
221 );
222}