1use std::fmt; 2 3use bytes::Bytes; 4use jev_protocol::retry::{self, Outcome}; 5use jev_protocol::{ 6 ApiError, ApiErrorKind, Json, ModelId, ProtocolError, Questions, RETRY_COUNT_HEADER, Response, request_bytes, 7 request_id, 8}; 9 10use crate::{ClientError, Event, HttpRequest, Observer, Runtime, Transport, TransportError, TransportErrorKind}; 11 12/// Resends of a request that never left, before it counts as a failure. 13/// A second dead connection in a row means something else is wrong. 14const MAX_REDIALS: u32 = 2; 15 16pub struct Client<T, R, O = ()> { 17 transport: T, 18 runtime: R, 19 observer: O, 20 model: ModelId, 21 authorization: http::HeaderValue, 22} 23 24/// A verified response and how it was obtained. 25#[derive(Debug)] 26pub struct Answered { 27 pub response: Response, 28 pub request_id: Option<String>, 29 pub attempts: u32, 30 /// The answering response's headers and body as they arrived, for a 31 /// caller that reads what the parse does not (rate-limit headers, a 32 /// field the protocol does not know yet). 33 pub headers: http::HeaderMap, 34 pub body: Bytes, 35} 36 37/// Why the last attempt failed, kept until retrying is ruled out. 38enum Failure { 39 Api(ApiError), 40 Transport(TransportError), 41 TimedOut, 42} 43 44impl fmt::Display for Failure { 45 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { 46 match self { 47 Failure::Api(error) => error.fmt(f), 48 Failure::Transport(error) => error.fmt(f), 49 Failure::TimedOut => f.write_str("attempt timed out"), 50 } 51 } 52} 53 54impl<T: Transport, R: Runtime> Client<T, R> { 55 pub fn new(transport: T, runtime: R, model: ModelId, api_key: &str) -> Result<Self, ClientError> { 56 let mut authorization = http::HeaderValue::try_from(format!("Bearer {api_key}")).map_err(|_| { 57 ClientError::Invalid(ProtocolError::Invalid("the API key has characters a header cannot carry".into())) 58 })?; 59 // Kept out of HPACK's dynamic table and out of Debug output. 60 authorization.set_sensitive(true); 61 Ok(Client { transport, runtime, observer: (), model, authorization }) 62 } 63 64 /// The same client, reporting to `observer`. 65 pub fn observed<O: Observer>(self, observer: O) -> Client<T, R, O> { 66 let Client { transport, runtime, model, authorization, .. } = self; 67 Client { transport, runtime, observer, model, authorization } 68 } 69} 70 71impl<T: Transport, R: Runtime, O: Observer> Client<T, R, O> { 72 pub fn model(&self) -> &ModelId { 73 &self.model 74 } 75 76 /// Asks `questions` about `state`, under the vendor SDKs' retry policy 77 /// (`jev_protocol::retry`), and verifies the answer. 78 pub async fn ask(&self, state: &Json, questions: &Questions) -> Result<Answered, ClientError> { 79 let body = Bytes::from(request_bytes(&self.model, state, questions).map_err(ClientError::Invalid)?); 80 let mut started = self.runtime.now(); 81 let mut attempts = 0; 82 let mut retries = 0; 83 let mut redials = 0; 84 let mut last_id: Option<String> = None; 85 loop { 86 let left = retry::BUDGET.saturating_sub(self.runtime.now() - started); 87 if left.is_zero() { 88 return Err(ClientError::TimedOut { attempts, request_id: last_id }); 89 } 90 let request = HttpRequest { headers: self.headers(retries), body: body.clone() }; 91 // Waiting for the account's rate limit spends none of the 92 // retry budget: nothing has been tried yet. 93 let queued = self.runtime.now(); 94 self.transport.admit(&request).await; 95 started += self.runtime.now() - queued; 96 attempts += 1; 97 self.observer.event(Event::Sending { retry: retries }); 98 let sent = self.runtime.timeout(retry::ATTEMPT_TIMEOUT.min(left), self.transport.send(request)).await; 99 100 let (failure, delay) = match sent { 101 Some(Ok(response)) => { 102 let id = request_id(&response.headers); 103 if id.is_some() { 104 last_id.clone_from(&id); 105 } 106 if response.status == 200 { 107 return match Response::parse(&self.model, questions, &response.body) { 108 Ok(parsed) => { 109 let usage = parsed.usage(); 110 self.observer.event(Event::Answered { request_id: id.as_deref(), attempts, usage }); 111 Ok(Answered { 112 response: parsed, 113 request_id: id, 114 attempts, 115 headers: response.headers, 116 body: response.body, 117 }) 118 } 119 Err(error) => Err(ClientError::Response { error, request_id: id }), 120 }; 121 } 122 let error = ApiError::from_response(response.status, &response.headers, &response.body); 123 // Too big for the context: splitting is the caller's 124 // move (the batch planner), never a resend. 125 if error.kind == ApiErrorKind::ContextOverflow { 126 return Err(ClientError::Api { error, attempts }); 127 } 128 let outcome = Outcome::Status { status: response.status, headers: &response.headers }; 129 let delay = retry::next_delay(&outcome, retries, self.runtime.jitter(), self.runtime.wall_clock()); 130 (Failure::Api(error), delay) 131 } 132 Some(Err(error)) => match error.kind { 133 TransportErrorKind::NotSent if redials < MAX_REDIALS => { 134 // Nothing reached the server: resend at once, as 135 // the same attempt, on a fresh connection. 136 self.observer.event(Event::Redialing { cause: &error }); 137 redials += 1; 138 attempts -= 1; 139 continue; 140 } 141 TransportErrorKind::Tls | TransportErrorKind::TooLarge | TransportErrorKind::Config => { 142 return Err(ClientError::Transport { error, attempts, request_id: last_id }); 143 } 144 _ => { 145 let delay = 146 retry::next_delay(&Outcome::Transport, retries, self.runtime.jitter(), self.runtime.wall_clock()); 147 (Failure::Transport(error), delay) 148 } 149 }, 150 None => { 151 self.transport.distrust_connection(); 152 let delay = 153 retry::next_delay(&Outcome::Transport, retries, self.runtime.jitter(), self.runtime.wall_clock()); 154 (Failure::TimedOut, delay) 155 } 156 }; 157 158 let fail = |failure: Failure, request_id: Option<String>| match failure { 159 Failure::Api(error) => ClientError::Api { error, attempts }, 160 Failure::Transport(error) => ClientError::Transport { error, attempts, request_id }, 161 Failure::TimedOut => ClientError::TimedOut { attempts, request_id }, 162 }; 163 let Some(delay) = delay else { return Err(fail(failure, last_id)) }; 164 if self.runtime.now() - started + delay >= retry::BUDGET { 165 return Err(fail(failure, last_id)); 166 } 167 self.observer.event(Event::Retrying { 168 attempt: attempts, 169 delay, 170 cause: &failure, 171 request_id: last_id.as_deref(), 172 }); 173 self.runtime.sleep(delay).await; 174 retries += 1; 175 } 176 } 177 178 fn headers(&self, retries: u32) -> http::HeaderMap { 179 let mut headers = http::HeaderMap::with_capacity(3); 180 headers.insert(http::header::AUTHORIZATION, self.authorization.clone()); 181 headers.insert(http::header::CONTENT_TYPE, http::HeaderValue::from_static("application/json")); 182 if retries > 0 { 183 // As both vendor SDKs do. 184 headers.insert(RETRY_COUNT_HEADER, http::HeaderValue::from(retries)); 185 } 186 headers 187 } 188} 189 190#[cfg(test)] 191mod tests;