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 rho_sdk::{ProviderError, ProviderErrorKind, Retryability};
9
10use crate::model::ModelError;
11
12/// Converts an application [`ModelError`] into a sanitized public [`ProviderError`].
13///
14/// HTTP response bodies and other transport payloads are omitted so credentials
15/// and provider-private content do not enter the SDK error contract.
16pub fn provider_error_from_model_error(error: ModelError) -> ProviderError {
17    match error {
18        ModelError::MissingApiKey
19        | ModelError::MissingCodexAuth
20        | ModelError::MissingAnthropicApiKey
21        | ModelError::MissingGithubCopilotAuth
22        | ModelError::MissingXaiAuth => ProviderError::new(
23            ProviderErrorKind::Authentication,
24            error.to_string(),
25            Retryability::Permanent,
26        ),
27        ModelError::Credentials(_) => ProviderError::new(
28            ProviderErrorKind::Authentication,
29            "credential store operation failed",
30            Retryability::Permanent,
31        ),
32        ModelError::Interrupted => ProviderError::interrupted("provider stream interrupted"),
33        ModelError::StreamIdleTimeout { timeout } => ProviderError::new(
34            ProviderErrorKind::Timeout,
35            format!(
36                "provider stream received no data for {timeout:?}; the connection may be stale"
37            ),
38            Retryability::Retryable,
39        ),
40        ModelError::StreamFailedAfterOutput { message: _ } => ProviderError::new(
41            ProviderErrorKind::InvalidResponse,
42            "provider stream failed after emitting output",
43            Retryability::Permanent,
44        ),
45        ModelError::InvalidResponse(_) => ProviderError::new(
46            ProviderErrorKind::InvalidResponse,
47            "provider returned an invalid response",
48            Retryability::Permanent,
49        ),
50        ModelError::UnsupportedReasoning {
51            provider,
52            model,
53            requested,
54        } => ProviderError::new(
55            ProviderErrorKind::Other,
56            format!(
57                "provider '{provider}' model '{model}' does not support reasoning level '{requested}'"
58            ),
59            Retryability::Permanent,
60        ),
61        ModelError::UnsupportedProvider(provider) => ProviderError::new(
62            ProviderErrorKind::Other,
63            format!("unsupported provider '{provider}'"),
64            Retryability::Permanent,
65        ),
66        ModelError::HttpStatus { status, body: _ } => {
67            let status_code = status.as_u16();
68            let (kind, retryability) = match status_code {
69                401 | 403 => (ProviderErrorKind::Authentication, Retryability::Permanent),
70                408 | 504 => (ProviderErrorKind::Timeout, Retryability::Retryable),
71                429 => (ProviderErrorKind::RateLimit, Retryability::Retryable),
72                500..=599 => (ProviderErrorKind::Unavailable, Retryability::Retryable),
73                _ => (ProviderErrorKind::Other, Retryability::Permanent),
74            };
75            ProviderError::new(kind, format!("HTTP {status_code}"), retryability)
76        }
77        ModelError::Request(_) => ProviderError::new(
78            ProviderErrorKind::Unavailable,
79            "provider request failed",
80            Retryability::Retryable,
81        ),
82        ModelError::Io(_) => ProviderError::new(
83            ProviderErrorKind::Other,
84            "provider I/O failed",
85            Retryability::Retryable,
86        ),
87    }
88}
89
90/// Implements [`rho_sdk::provider::ModelProvider`] for an application transport
91/// that already exposes inherent `model_identity`, `complete_turn`, and
92/// `stream_turn` methods.
93///
94/// Streaming uses a bounded callback bridge. A callback burst that fills the
95/// bridge is interrupted rather than buffered without bound.
96#[macro_export]
97macro_rules! impl_sdk_model_provider {
98    ($provider:ty) => {
99        impl ::rho_sdk::provider::ModelProvider for $provider {
100            fn identity(&self) -> ::rho_sdk::model::ModelIdentity {
101                self.model_identity()
102            }
103
104            fn send_turn<'a>(
105                &'a self,
106                request: ::rho_sdk::model::ModelRequest<'a>,
107            ) -> ::rho_sdk::provider::ProviderFuture<'a> {
108                ::std::boxed::Box::pin(async move {
109                    self.complete_turn(request)
110                        .await
111                        .map_err($crate::providers::sdk_contract::provider_error_from_model_error)
112                })
113            }
114
115            fn send_turn_stream<'a>(
116                &'a self,
117                request: ::rho_sdk::model::ModelRequest<'a>,
118                events: ::rho_sdk::provider::ProviderEventSender,
119            ) -> ::rho_sdk::provider::ProviderFuture<'a> {
120                ::std::boxed::Box::pin(async move {
121                    let (event_tx, mut event_rx) =
122                        ::tokio::sync::mpsc::channel(events.capacity());
123                    let mut on_event = move |event| {
124                        event_tx
125                            .try_send(event)
126                            .map_err(|_| $crate::model::ModelError::Interrupted)
127                    };
128                    let mut provider = ::std::pin::pin!(self.stream_turn(request, &mut on_event));
129                    loop {
130                        ::tokio::select! {
131                            biased;
132                            event = event_rx.recv() => {
133                                if let Some(event) = event {
134                                    events.send(event).await?;
135                                }
136                            }
137                            result = &mut provider => {
138                                while let Ok(event) = event_rx.try_recv() {
139                                    events.send(event).await?;
140                                }
141                                return result.map_err(
142                                    $crate::providers::sdk_contract::provider_error_from_model_error,
143                                );
144                            }
145                        }
146                    }
147                })
148            }
149        }
150    };
151}
152
153#[cfg(test)]
154#[path = "sdk_contract_tests.rs"]
155mod tests;