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::MissingMoonshotApiKey
34        | ModelError::MissingOpenRouterApiKey
35        | ModelError::MissingKimiAuth
36        | ModelError::MissingXaiApiKey
37        | ModelError::MissingXaiAuth => ProviderError::new(
38            ProviderErrorKind::Authentication,
39            error.to_string(),
40            Retryability::Permanent,
41        ),
42        ModelError::Credentials(_) => ProviderError::new(
43            ProviderErrorKind::Authentication,
44            "credential store operation failed",
45            Retryability::Permanent,
46        ),
47        ModelError::Interrupted => ProviderError::interrupted("provider stream interrupted"),
48        ModelError::StreamIdleTimeout { timeout } => ProviderError::new(
49            ProviderErrorKind::Timeout,
50            format!(
51                "provider stream received no data for {timeout:?}; the connection may be stale"
52            ),
53            Retryability::Retryable,
54        ),
55        ModelError::StreamFailedAfterOutput { message } => ProviderError::new(
56            ProviderErrorKind::InvalidResponse,
57            "provider stream failed after emitting output",
58            Retryability::Permanent,
59        )
60        .with_diagnostic(sanitize_diagnostic(&message)),
61        ModelError::InvalidResponse(details) => ProviderError::new(
62            ProviderErrorKind::InvalidResponse,
63            "provider returned an invalid response",
64            Retryability::Permanent,
65        )
66        .with_diagnostic(sanitize_diagnostic(&details)),
67        ModelError::UnsupportedReasoning {
68            provider,
69            model,
70            requested,
71        } => ProviderError::new(
72            ProviderErrorKind::Other,
73            format!(
74                "provider '{provider}' model '{model}' does not support reasoning level '{requested}'"
75            ),
76            Retryability::Permanent,
77        ),
78        ModelError::UnsupportedProvider(provider) => ProviderError::new(
79            ProviderErrorKind::Other,
80            format!("unsupported provider '{provider}'"),
81            Retryability::Permanent,
82        ),
83        ModelError::HttpStatus { status, body } => {
84            let status_code = status.as_u16();
85            let (kind, retryability) = match status_code {
86                401 | 403 => (ProviderErrorKind::Authentication, Retryability::Permanent),
87                408 | 504 => (ProviderErrorKind::Timeout, Retryability::Retryable),
88                429 => (ProviderErrorKind::RateLimit, Retryability::Retryable),
89                500..=599 => (ProviderErrorKind::Unavailable, Retryability::Retryable),
90                _ => (ProviderErrorKind::Other, Retryability::Permanent),
91            };
92            let error = ProviderError::new(kind, format!("HTTP {status_code}"), retryability);
93            if body.is_empty() {
94                error
95            } else {
96                error.with_diagnostic(sanitize_diagnostic(&body))
97            }
98        }
99        ModelError::Request(_) => ProviderError::new(
100            ProviderErrorKind::Unavailable,
101            "provider request failed",
102            Retryability::Retryable,
103        ),
104        ModelError::Io(_) => ProviderError::new(
105            ProviderErrorKind::Other,
106            "provider I/O failed",
107            Retryability::Retryable,
108        ),
109    }
110}
111
112fn sanitize_diagnostic(value: &str) -> String {
113    const MAX_BYTES: usize = crate::provider_backend::http_error::MAX_ERROR_BODY_BYTES;
114
115    let mut diagnostic = String::new();
116    let mut truncated = false;
117    for character in value.chars() {
118        let escaped = match character {
119            '\n' | '\t' => character.to_string(),
120            character if character.is_control() => character.escape_default().to_string(),
121            character => character.to_string(),
122        };
123        if diagnostic.len() + escaped.len() > MAX_BYTES {
124            truncated = true;
125            break;
126        }
127        diagnostic.push_str(&escaped);
128    }
129    if truncated {
130        diagnostic.push_str("\n[diagnostic truncated]");
131    }
132    diagnostic
133}
134
135/// Shared queue used by [`callback_event_sink`] and [`drive_callback_stream`].
136pub type CallbackEventQueue = Arc<Mutex<VecDeque<ModelEvent>>>;
137
138/// Builds the synchronous callback used by application stream transports.
139///
140/// Events are buffered temporarily because the callback cannot await. The
141/// companion [`drive_callback_stream`] loop drains that buffer through the
142/// SDK's bounded event sender before polling the provider again.
143pub fn callback_event_sink(
144    cancellation: CancellationToken,
145    pending: CallbackEventQueue,
146    notify: Arc<Notify>,
147) -> impl FnMut(ModelEvent) -> Result<(), ModelError> + Send {
148    move |event| {
149        if cancellation.is_cancelled() {
150            return Err(ModelError::Interrupted);
151        }
152        pending
153            .lock()
154            .unwrap_or_else(|poisoned| poisoned.into_inner())
155            .push_back(event);
156        notify.notify_one();
157        Ok(())
158    }
159}
160
161/// Drains buffered callback events through the bounded SDK event channel and
162/// drives the provider future with host backpressure across awaits.
163pub async fn drive_callback_stream<Fut>(
164    cancellation: CancellationToken,
165    events: ProviderEventSender,
166    pending: CallbackEventQueue,
167    notify: Arc<Notify>,
168    provider: Fut,
169) -> Result<ModelResponse, ProviderError>
170where
171    Fut: Future<Output = Result<ModelResponse, ModelError>>,
172{
173    let mut provider = std::pin::pin!(provider);
174    let mut provider_result: Option<Result<ModelResponse, ModelError>> = None;
175
176    loop {
177        loop {
178            let next = pending
179                .lock()
180                .unwrap_or_else(|poisoned| poisoned.into_inner())
181                .pop_front();
182            let Some(event) = next else {
183                break;
184            };
185            if cancellation.is_cancelled() {
186                return Err(ProviderError::interrupted("provider stream interrupted"));
187            }
188            events.send(event).await?;
189        }
190
191        if let Some(result) = provider_result.take() {
192            return result.map_err(provider_error_from_model_error);
193        }
194
195        let notified = notify.notified();
196        let has_pending = !pending
197            .lock()
198            .unwrap_or_else(|poisoned| poisoned.into_inner())
199            .is_empty();
200        if has_pending {
201            continue;
202        }
203
204        tokio::select! {
205            biased;
206            () = notified => {}
207            () = cancellation.cancelled() => {
208                return Err(ProviderError::interrupted("provider stream interrupted"));
209            }
210            result = &mut provider => {
211                provider_result = Some(result);
212            }
213        }
214    }
215}
216
217/// Implements [`rho_sdk::provider::ModelProvider`] for an application transport
218/// that already exposes inherent `model_identity`, `complete_turn`, and
219/// `stream_turn` methods.
220///
221/// Streaming buffers same-poll callback bursts, then applies the SDK event
222/// channel's async backpressure before polling the provider again.
223#[macro_export]
224macro_rules! impl_sdk_model_provider {
225    ($provider:ty) => {
226        impl ::rho_sdk::provider::ModelProvider for $provider {
227            fn identity(&self) -> ::rho_sdk::model::ModelIdentity {
228                self.model_identity()
229            }
230
231            fn send_turn<'a>(
232                &'a self,
233                request: ::rho_sdk::model::ModelRequest<'a>,
234            ) -> ::rho_sdk::provider::ProviderFuture<'a> {
235                ::std::boxed::Box::pin(async move {
236                    self.complete_turn(request)
237                        .await
238                        .map_err($crate::providers::sdk_contract::provider_error_from_model_error)
239                })
240            }
241
242            fn send_turn_stream<'a>(
243                &'a self,
244                request: ::rho_sdk::model::ModelRequest<'a>,
245                events: ::rho_sdk::provider::ProviderEventSender,
246            ) -> ::rho_sdk::provider::ProviderFuture<'a> {
247                ::std::boxed::Box::pin(async move {
248                    let cancellation = request.cancellation.clone();
249                    let pending = ::std::sync::Arc::new(::std::sync::Mutex::new(
250                        ::std::collections::VecDeque::new(),
251                    ));
252                    let notify = ::std::sync::Arc::new(::tokio::sync::Notify::new());
253                    let mut on_event = $crate::providers::sdk_contract::callback_event_sink(
254                        cancellation.clone(),
255                        ::std::sync::Arc::clone(&pending),
256                        ::std::sync::Arc::clone(&notify),
257                    );
258                    let provider = self.stream_turn(request, &mut on_event);
259                    $crate::providers::sdk_contract::drive_callback_stream(
260                        cancellation,
261                        events,
262                        pending,
263                        notify,
264                        provider,
265                    )
266                    .await
267                })
268            }
269        }
270    };
271}
272
273#[cfg(test)]
274#[path = "sdk_contract_tests.rs"]
275mod tests;