Skip to main content

rho_coding_agent/providers/
sdk_contract.rs

1//! Shared helpers for exposing application transports through the public SDK
2//! provider contract.
3//!
4//! Built-in providers implement [`rho_sdk::provider::ModelProvider`] directly.
5//! Callback-based stream transports remain an internal detail and are bridged
6//! here into the SDK's bounded async event sender.
7
8use std::{
9    collections::VecDeque,
10    future::Future,
11    sync::{Arc, Mutex},
12};
13
14use rho_sdk::{
15    model::{ModelEvent, ModelResponse},
16    provider::ProviderEventSender,
17    CancellationToken, ProviderError, ProviderErrorKind, Retryability,
18};
19use tokio::sync::Notify;
20
21use crate::model::ModelError;
22
23/// Converts an application [`ModelError`] into a sanitized public [`ProviderError`].
24///
25/// HTTP response bodies and other transport payloads are omitted so credentials
26/// and provider-private content do not enter the SDK error contract.
27pub fn provider_error_from_model_error(error: ModelError) -> ProviderError {
28    match error {
29        ModelError::MissingApiKey
30        | ModelError::MissingCodexAuth
31        | ModelError::MissingAnthropicApiKey
32        | ModelError::MissingGithubCopilotAuth
33        | ModelError::MissingXaiAuth => ProviderError::new(
34            ProviderErrorKind::Authentication,
35            error.to_string(),
36            Retryability::Permanent,
37        ),
38        ModelError::Credentials(_) => ProviderError::new(
39            ProviderErrorKind::Authentication,
40            "credential store operation failed",
41            Retryability::Permanent,
42        ),
43        ModelError::Interrupted => ProviderError::interrupted("provider stream interrupted"),
44        ModelError::StreamIdleTimeout { timeout } => ProviderError::new(
45            ProviderErrorKind::Timeout,
46            format!(
47                "provider stream received no data for {timeout:?}; the connection may be stale"
48            ),
49            Retryability::Retryable,
50        ),
51        ModelError::StreamFailedAfterOutput { message: _ } => ProviderError::new(
52            ProviderErrorKind::InvalidResponse,
53            "provider stream failed after emitting output",
54            Retryability::Permanent,
55        ),
56        ModelError::InvalidResponse(_) => ProviderError::new(
57            ProviderErrorKind::InvalidResponse,
58            "provider returned an invalid response",
59            Retryability::Permanent,
60        ),
61        ModelError::UnsupportedReasoning {
62            provider,
63            model,
64            requested,
65        } => ProviderError::new(
66            ProviderErrorKind::Other,
67            format!(
68                "provider '{provider}' model '{model}' does not support reasoning level '{requested}'"
69            ),
70            Retryability::Permanent,
71        ),
72        ModelError::UnsupportedProvider(provider) => ProviderError::new(
73            ProviderErrorKind::Other,
74            format!("unsupported provider '{provider}'"),
75            Retryability::Permanent,
76        ),
77        ModelError::HttpStatus { status, body: _ } => {
78            let status_code = status.as_u16();
79            let (kind, retryability) = match status_code {
80                401 | 403 => (ProviderErrorKind::Authentication, Retryability::Permanent),
81                408 | 504 => (ProviderErrorKind::Timeout, Retryability::Retryable),
82                429 => (ProviderErrorKind::RateLimit, Retryability::Retryable),
83                500..=599 => (ProviderErrorKind::Unavailable, Retryability::Retryable),
84                _ => (ProviderErrorKind::Other, Retryability::Permanent),
85            };
86            ProviderError::new(kind, format!("HTTP {status_code}"), retryability)
87        }
88        ModelError::Request(_) => ProviderError::new(
89            ProviderErrorKind::Unavailable,
90            "provider request failed",
91            Retryability::Retryable,
92        ),
93        ModelError::Io(_) => ProviderError::new(
94            ProviderErrorKind::Other,
95            "provider I/O failed",
96            Retryability::Retryable,
97        ),
98    }
99}
100
101/// Shared queue used by [`callback_event_sink`] and [`drive_callback_stream`].
102pub type CallbackEventQueue = Arc<Mutex<VecDeque<ModelEvent>>>;
103
104/// Builds the synchronous callback used by application stream transports.
105///
106/// Events are buffered temporarily because the callback cannot await. The
107/// companion [`drive_callback_stream`] loop drains that buffer through the
108/// SDK's bounded event sender before polling the provider again.
109pub fn callback_event_sink(
110    cancellation: CancellationToken,
111    pending: CallbackEventQueue,
112    notify: Arc<Notify>,
113) -> impl FnMut(ModelEvent) -> Result<(), ModelError> + Send {
114    move |event| {
115        if cancellation.is_cancelled() {
116            return Err(ModelError::Interrupted);
117        }
118        pending
119            .lock()
120            .unwrap_or_else(|poisoned| poisoned.into_inner())
121            .push_back(event);
122        notify.notify_one();
123        Ok(())
124    }
125}
126
127/// Drains buffered callback events through the bounded SDK event channel and
128/// drives the provider future with host backpressure across awaits.
129pub async fn drive_callback_stream<Fut>(
130    cancellation: CancellationToken,
131    events: ProviderEventSender,
132    pending: CallbackEventQueue,
133    notify: Arc<Notify>,
134    provider: Fut,
135) -> Result<ModelResponse, ProviderError>
136where
137    Fut: Future<Output = Result<ModelResponse, ModelError>>,
138{
139    let mut provider = std::pin::pin!(provider);
140    let mut provider_result: Option<Result<ModelResponse, ModelError>> = None;
141
142    loop {
143        loop {
144            let next = pending
145                .lock()
146                .unwrap_or_else(|poisoned| poisoned.into_inner())
147                .pop_front();
148            let Some(event) = next else {
149                break;
150            };
151            if cancellation.is_cancelled() {
152                return Err(ProviderError::interrupted("provider stream interrupted"));
153            }
154            events.send(event).await?;
155        }
156
157        if let Some(result) = provider_result.take() {
158            return result.map_err(provider_error_from_model_error);
159        }
160
161        let notified = notify.notified();
162        let has_pending = !pending
163            .lock()
164            .unwrap_or_else(|poisoned| poisoned.into_inner())
165            .is_empty();
166        if has_pending {
167            continue;
168        }
169
170        tokio::select! {
171            biased;
172            () = notified => {}
173            () = cancellation.cancelled() => {
174                return Err(ProviderError::interrupted("provider stream interrupted"));
175            }
176            result = &mut provider => {
177                provider_result = Some(result);
178            }
179        }
180    }
181}
182
183/// Implements [`rho_sdk::provider::ModelProvider`] for an application transport
184/// that already exposes inherent `model_identity`, `complete_turn`, and
185/// `stream_turn` methods.
186///
187/// Streaming buffers same-poll callback bursts, then applies the SDK event
188/// channel's async backpressure before polling the provider again.
189#[macro_export]
190macro_rules! impl_sdk_model_provider {
191    ($provider:ty) => {
192        impl ::rho_sdk::provider::ModelProvider for $provider {
193            fn identity(&self) -> ::rho_sdk::model::ModelIdentity {
194                self.model_identity()
195            }
196
197            fn send_turn<'a>(
198                &'a self,
199                request: ::rho_sdk::model::ModelRequest<'a>,
200            ) -> ::rho_sdk::provider::ProviderFuture<'a> {
201                ::std::boxed::Box::pin(async move {
202                    self.complete_turn(request)
203                        .await
204                        .map_err($crate::providers::sdk_contract::provider_error_from_model_error)
205                })
206            }
207
208            fn send_turn_stream<'a>(
209                &'a self,
210                request: ::rho_sdk::model::ModelRequest<'a>,
211                events: ::rho_sdk::provider::ProviderEventSender,
212            ) -> ::rho_sdk::provider::ProviderFuture<'a> {
213                ::std::boxed::Box::pin(async move {
214                    let cancellation = request.cancellation.clone();
215                    let pending = ::std::sync::Arc::new(::std::sync::Mutex::new(
216                        ::std::collections::VecDeque::new(),
217                    ));
218                    let notify = ::std::sync::Arc::new(::tokio::sync::Notify::new());
219                    let mut on_event = $crate::providers::sdk_contract::callback_event_sink(
220                        cancellation.clone(),
221                        ::std::sync::Arc::clone(&pending),
222                        ::std::sync::Arc::clone(&notify),
223                    );
224                    let provider = self.stream_turn(request, &mut on_event);
225                    $crate::providers::sdk_contract::drive_callback_stream(
226                        cancellation,
227                        events,
228                        pending,
229                        notify,
230                        provider,
231                    )
232                    .await
233                })
234            }
235        }
236    };
237}
238
239#[cfg(test)]
240#[path = "sdk_contract_tests.rs"]
241mod tests;