rho_coding_agent/providers/
sdk_contract.rs1use rho_sdk::{ProviderError, ProviderErrorKind, Retryability};
9
10use crate::model::ModelError;
11
12pub 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#[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;