Skip to main content

starweaver_model/wrappers/
concurrency.rs

1//! Concurrency-limited model wrapper.
2
3use 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
21/// Model wrapper that limits concurrent calls with a shared semaphore.
22pub 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    /// Create a concurrency-limited wrapper.
35    #[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    /// Create a wrapper using an existing shared semaphore.
46    #[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    /// Return configured max concurrency.
56    #[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}