rho_coding_agent/providers/
sdk_contract.rs1use 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
23pub 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
101pub type CallbackEventQueue = Arc<Mutex<VecDeque<ModelEvent>>>;
103
104pub 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
127pub 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#[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(¬ify),
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;