1use std::collections::BTreeMap;
5use std::fmt;
6use std::sync::Arc;
7use std::time::{Duration, Instant};
8
9use serde_json::Value;
10
11use crate::answers::{decode_models, ModelMetadata, SystemOneResponse};
12use crate::call_options::{CallOptions, ResolvedCall};
13use crate::config::{
14 resolve_api_key, resolve_base_url, resolve_default_model, validate_timeout, LogLevel,
15};
16use crate::errors::{api_error_message, request_id_of, ApiError, ApiErrorKind, Error};
17use crate::logging;
18use crate::request::SystemOneRequest;
19use crate::retry::RetryPolicy;
20
21pub(crate) struct RawResponse {
23 pub status: u16,
25 pub body: String,
27 pub request_id: Option<String>,
29 pub endpoint: String,
31}
32
33#[derive(Debug, Clone, Default)]
56pub struct RequestOptions {
57 pub(crate) options: CallOptions,
58}
59
60impl RequestOptions {
61 pub fn new() -> Self {
63 Self::default()
64 }
65
66 pub fn timeout(mut self, timeout: Duration) -> Self {
68 self.options.timeout = Some(timeout);
69 self
70 }
71
72 pub fn retry_policy(mut self, policy: RetryPolicy) -> Self {
74 self.options.retry_policy = Some(policy);
75 self
76 }
77
78 pub fn header(mut self, name: impl Into<String>, value: impl Into<String>) -> Self {
80 self.options.headers.insert(name.into(), value.into());
81 self
82 }
83
84 pub fn extra_body(mut self, name: impl Into<String>, value: Value) -> Self {
89 self.options
90 .extra_body
91 .get_or_insert_with(serde_json::Map::new)
92 .insert(name.into(), value);
93 self
94 }
95}
96
97struct Inner {
99 api_key: String,
100 base_url: String,
101 default_model: String,
102 timeout: Duration,
103 retry_policy: RetryPolicy,
104 default_headers: BTreeMap<String, String>,
105 log_level: LogLevel,
106 http: reqwest::Client,
107}
108
109#[derive(Clone)]
128pub struct Client {
129 inner: Arc<Inner>,
130}
131
132impl fmt::Debug for Client {
133 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
134 f.debug_struct("Client")
136 .field("base_url", &self.inner.base_url)
137 .field("default_model", &self.inner.default_model)
138 .field("timeout", &self.inner.timeout)
139 .field("retry_policy", &self.inner.retry_policy)
140 .field("default_headers", &self.inner.default_headers)
141 .field("log_level", &self.inner.log_level)
142 .finish_non_exhaustive()
143 }
144}
145
146#[derive(Default)]
154pub struct ClientBuilder {
155 api_key: Option<String>,
156 base_url: Option<String>,
157 default_model: Option<String>,
158 timeout: Option<Duration>,
159 retry_policy: Option<RetryPolicy>,
160 default_headers: BTreeMap<String, String>,
161 log_level: Option<LogLevel>,
162 http_client: Option<reqwest::Client>,
163}
164
165impl fmt::Debug for ClientBuilder {
166 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
167 f.debug_struct("ClientBuilder")
169 .field("base_url", &self.base_url)
170 .field("default_model", &self.default_model)
171 .field("timeout", &self.timeout)
172 .field("retry_policy", &self.retry_policy)
173 .field("default_headers", &self.default_headers)
174 .field("log_level", &self.log_level)
175 .field(
176 "http_client",
177 &self.http_client.as_ref().map(|_| "reqwest::Client"),
178 )
179 .finish()
180 }
181}
182
183impl ClientBuilder {
184 pub fn api_key(mut self, api_key: impl Into<String>) -> Self {
190 self.api_key = Some(api_key.into());
191 self
192 }
193
194 pub fn base_url(mut self, base_url: impl Into<String>) -> Self {
200 self.base_url = Some(base_url.into());
201 self
202 }
203
204 pub fn default_model(mut self, model: impl Into<String>) -> Self {
207 self.default_model = Some(model.into());
208 self
209 }
210
211 pub fn timeout(mut self, timeout: Duration) -> Self {
214 self.timeout = Some(timeout);
215 self
216 }
217
218 pub fn retry_policy(mut self, policy: RetryPolicy) -> Self {
220 self.retry_policy = Some(policy);
221 self
222 }
223
224 pub fn default_header(mut self, name: impl Into<String>, value: impl Into<String>) -> Self {
230 self.default_headers.insert(name.into(), value.into());
231 self
232 }
233
234 pub fn log_level(mut self, level: LogLevel) -> Self {
236 self.log_level = Some(level);
237 self
238 }
239
240 pub fn http_client(mut self, http: reqwest::Client) -> Self {
242 self.http_client = Some(http);
243 self
244 }
245
246 pub fn build(self) -> Result<Client, Error> {
248 let api_key = resolve_api_key(self.api_key.as_deref())?;
249 let base_url = resolve_base_url(self.base_url.as_deref());
250 validate_base_url(&base_url)?;
251 let default_model = resolve_default_model(self.default_model.as_deref());
252 let timeout = match self.timeout {
253 Some(timeout) => validate_timeout(timeout)?,
254 None => crate::DEFAULT_TIMEOUT,
255 };
256 let retry_policy = self.retry_policy.unwrap_or_default();
257 retry_policy.validate()?;
258 let log_level = LogLevel::resolve(self.log_level);
259 let http = self.http_client.unwrap_or_default();
260 Ok(Client {
261 inner: Arc::new(Inner {
262 api_key,
263 base_url,
264 default_model,
265 timeout,
266 retry_policy,
267 default_headers: self.default_headers,
268 log_level,
269 http,
270 }),
271 })
272 }
273}
274
275impl Client {
276 pub fn builder() -> ClientBuilder {
278 ClientBuilder::default()
279 }
280
281 pub fn from_env() -> Result<Self, Error> {
285 Self::builder().build()
286 }
287
288 pub fn base_url(&self) -> &str {
290 &self.inner.base_url
291 }
292
293 pub fn default_model(&self) -> &str {
295 &self.inner.default_model
296 }
297
298 pub fn timeout(&self) -> Duration {
300 self.inner.timeout
301 }
302
303 pub fn retry_policy(&self) -> &RetryPolicy {
305 &self.inner.retry_policy
306 }
307
308 pub fn log_level(&self) -> LogLevel {
310 self.inner.log_level
311 }
312
313 pub async fn system_one(&self, request: SystemOneRequest) -> Result<SystemOneResponse, Error> {
319 self.system_one_with(request, RequestOptions::new()).await
320 }
321
322 pub async fn system_one_with(
327 &self,
328 request: SystemOneRequest,
329 options: RequestOptions,
330 ) -> Result<SystemOneResponse, Error> {
331 request.validate()?;
332 let call = self.resolve_call(options.options)?;
333 let model = request
334 .model
335 .clone()
336 .unwrap_or_else(|| self.inner.default_model.clone());
337 let mut body = serde_json::Map::new();
338 body.insert("state".into(), request.state.clone());
339 body.insert("model".into(), Value::String(model));
340 body.insert(
341 "questions".into(),
342 serde_json::to_value(&request.questions).map_err(|error| {
343 Error::InvalidRequest(format!("questions failed to serialize: {error}"))
344 })?,
345 );
346 if let Some(extra) = call.extra_body.clone() {
347 for (name, value) in extra {
348 body.insert(name, value);
349 }
350 }
351 let path = "/v1/systemone";
352 let url = format!("{}{}", self.inner.base_url, path);
353 let raw = self
354 .execute(
355 reqwest::Method::POST,
356 &url,
357 path,
358 Some(&Value::Object(body)),
359 &call,
360 )
361 .await?;
362 let response = SystemOneResponse::decode(&raw)?;
363 let overridden = call
366 .extra_body
367 .as_ref()
368 .is_some_and(|extra| extra.iter().any(|(name, _)| name == "questions"));
369 if !overridden {
370 crate::answers::check_complete(&raw, &response.answers, &request.questions)?;
371 }
372 Ok(response)
373 }
374
375 pub async fn list_models(&self) -> Result<Vec<ModelMetadata>, Error> {
377 self.list_models_with(RequestOptions::new()).await
378 }
379
380 pub async fn list_models_with(
385 &self,
386 options: RequestOptions,
387 ) -> Result<Vec<ModelMetadata>, Error> {
388 let call = self.resolve_call(options.options)?;
389 let path = "/v1/models";
390 let url = format!("{}{}", self.inner.base_url, path);
391 let raw = self
392 .execute(reqwest::Method::GET, &url, path, None, &call)
393 .await?;
394 decode_models(&raw)
395 }
396
397 fn resolve_call(&self, options: CallOptions) -> Result<ResolvedCall, Error> {
402 let per_call = |error: Error| match error {
403 Error::Config(message) => Error::InvalidRequest(format!("per-call option: {message}")),
404 other => other,
405 };
406 let timeout = match options.timeout {
407 Some(timeout) => crate::config::validate_timeout(timeout).map_err(per_call)?,
408 None => self.inner.timeout,
409 };
410 let retry_policy = match options.retry_policy {
411 Some(policy) => {
412 policy.validate().map_err(per_call)?;
413 policy
414 }
415 None => self.inner.retry_policy.clone(),
416 };
417 Ok(ResolvedCall {
418 timeout,
419 retry_policy,
420 headers: options.headers,
421 extra_body: options.extra_body,
422 })
423 }
424
425 async fn execute(
431 &self,
432 method: reqwest::Method,
433 url: &str,
434 path: &str,
435 body: Option<&Value>,
436 call: &ResolvedCall,
437 ) -> Result<RawResponse, Error> {
438 let endpoint = format!("{method} {url}");
439 let mut headers: BTreeMap<String, String> = BTreeMap::new();
446 for (name, value) in self.inner.default_headers.iter().chain(call.headers.iter()) {
447 headers.insert(name.to_ascii_lowercase(), value.clone());
448 }
449 headers.remove("x-typesafe-retry-count");
450 if body.is_none() {
451 headers.remove("content-type");
452 }
453 headers.insert(
454 "authorization".into(),
455 format!("Bearer {}", self.inner.api_key),
456 );
457 headers.insert("accept".into(), "application/json".into());
458 if body.is_some() {
459 headers.insert("content-type".into(), "application/json".into());
460 }
461 headers.insert("user-agent".into(), sdk_version_header());
462 headers.insert("x-typesafe-sdk".into(), sdk_version_header());
463 headers.insert("x-typesafe-runtime".into(), runtime_header());
464
465 let policy = call.retry_policy.clone();
466 let total_budget = policy.total_budget.filter(|budget| !budget.is_zero());
467 let call_start = Instant::now();
468 let mut last_error: Option<Error> = None;
469
470 for retry_number in 0..=policy.max_retries {
471 let mut attempt_headers = headers.clone();
472 if retry_number > 0 {
473 attempt_headers.insert("x-typesafe-retry-count".into(), retry_number.to_string());
474 }
475 let reqwest_headers = build_header_map(&attempt_headers)?;
476 if self.inner.log_level >= LogLevel::Debug {
477 logging::log_request_debug(method.as_str(), path, &reqwest_headers, body);
478 }
479
480 let started = Instant::now();
481 let mut attempt = self
482 .inner
483 .http
484 .request(method.clone(), url)
485 .headers(reqwest_headers.clone())
486 .timeout(call.timeout);
487 if let Some(body) = body {
488 attempt = attempt.body(body.to_string());
489 }
490 let result = attempt.send().await;
491 let response = match result {
492 Ok(response) => response,
493 Err(error) => {
494 let error: Error = if error.is_timeout() {
495 Error::Timeout {
496 timeout: call.timeout,
497 }
498 } else {
499 Error::Connection {
502 message: format!("Connection error: {error}"),
503 source: Some(Box::new(error)),
504 }
505 };
506 if self.inner.log_level >= LogLevel::Info {
507 logging::log_attempt_error(method.as_str(), path, &error.to_string());
508 }
509 last_error = Some(error);
510 if !should_retry(&policy, &last_error, retry_number) {
512 return Err(last_error.unwrap());
513 }
514 let delay = policy.delay_for_retry(retry_number, None);
515 if !delay_within_budget(&delay, &total_budget, call_start) {
516 return Err(last_error.unwrap());
518 }
519 if self.inner.log_level >= LogLevel::Info {
520 logging::log_retry_scheduled(
521 method.as_str(),
522 path,
523 delay,
524 retry_number + 1,
525 &last_error.as_ref().unwrap().to_string(),
526 );
527 }
528 tokio::time::sleep(delay).await;
529 continue;
530 }
531 };
532
533 let status = response.status().as_u16();
534 let response_headers = response.headers().clone();
535 let request_id = request_id_of(&response_headers);
536 let text = match response.text().await {
537 Ok(text) => text,
538 Err(error) => {
541 let error = if error.is_timeout() {
542 Error::Timeout {
543 timeout: call.timeout,
544 }
545 } else {
546 Error::Connection {
547 message: format!("Connection error: {error}"),
548 source: Some(Box::new(error)),
549 }
550 };
551 if self.inner.log_level >= LogLevel::Info {
552 logging::log_attempt_error(method.as_str(), path, &error.to_string());
553 }
554 last_error = Some(error);
555 if !should_retry(&policy, &last_error, retry_number) {
556 return Err(last_error.unwrap());
557 }
558 let delay = policy.delay_for_retry(retry_number, None);
559 if !delay_within_budget(&delay, &total_budget, call_start) {
560 return Err(last_error.unwrap());
561 }
562 if self.inner.log_level >= LogLevel::Info {
563 logging::log_retry_scheduled(
564 method.as_str(),
565 path,
566 delay,
567 retry_number + 1,
568 &last_error.as_ref().unwrap().to_string(),
569 );
570 }
571 tokio::time::sleep(delay).await;
572 continue;
573 }
574 };
575 let duration = started.elapsed();
576 if self.inner.log_level >= LogLevel::Info {
577 logging::log_attempt_info(
578 method.as_str(),
579 path,
580 status,
581 duration,
582 request_id.as_deref(),
583 );
584 }
585 if self.inner.log_level >= LogLevel::Debug {
586 logging::log_response_debug(
587 method.as_str(),
588 path,
589 status,
590 &response_headers,
591 &text,
592 );
593 }
594
595 if (200..300).contains(&status) {
596 return Ok(RawResponse {
597 status,
598 body: text,
599 request_id,
600 endpoint,
601 });
602 }
603
604 let parsed_body = parse_body(&text);
606 let retry_after = policy.parse_retry_after(&response_headers);
609 let api_error = ApiError {
610 status,
611 kind: ApiErrorKind::from_status(status),
612 message: api_error_message(&parsed_body),
613 body: parsed_body,
614 headers: response_headers,
615 request_id,
616 endpoint: endpoint.clone(),
617 retry_after,
618 };
619 let error = Error::from(api_error);
620 last_error = Some(error);
621
622 if !should_retry(&policy, &last_error, retry_number) {
623 return Err(last_error.unwrap());
624 }
625 let delay =
626 policy.delay_for_retry(retry_number, last_error.as_ref().and_then(Error::as_api));
627 if !delay_within_budget(&delay, &total_budget, call_start) {
628 return Err(last_error.unwrap());
631 }
632 if self.inner.log_level >= LogLevel::Info {
633 logging::log_retry_scheduled(
634 method.as_str(),
635 path,
636 delay,
637 retry_number + 1,
638 &last_error.as_ref().unwrap().to_string(),
639 );
640 }
641 tokio::time::sleep(delay).await;
642 }
643
644 Err(last_error.unwrap_or_else(|| Error::Connection {
645 message: "Connection error: exhausted retries.".into(),
646 source: None,
647 }))
648 }
649}
650
651fn should_retry(policy: &RetryPolicy, last_error: &Option<Error>, retry_number: u32) -> bool {
653 if retry_number >= policy.max_retries {
654 return false;
655 }
656 last_error
657 .as_ref()
658 .map(|error| policy.is_retryable(error))
659 .unwrap_or(false)
660}
661
662fn delay_within_budget(
664 delay: &Duration,
665 total_budget: &Option<Duration>,
666 call_start: Instant,
667) -> bool {
668 match total_budget {
669 None => true,
670 Some(budget) => {
671 let elapsed = call_start.elapsed();
672 elapsed.saturating_add(*delay) < *budget
674 }
675 }
676}
677
678fn parse_body(text: &str) -> Option<Value> {
681 if text.is_empty() {
682 return None;
683 }
684 match serde_json::from_str(text) {
685 Ok(value) => Some(value),
686 Err(_) => Some(Value::String(text.to_owned())),
687 }
688}
689
690fn build_header_map(
691 headers: &BTreeMap<String, String>,
692) -> Result<reqwest::header::HeaderMap, Error> {
693 let mut map = reqwest::header::HeaderMap::new();
694 for (name, value) in headers {
695 let name: reqwest::header::HeaderName = name
696 .parse()
697 .map_err(|error| Error::Config(format!("invalid header name {name:?}: {error}")))?;
698 let value = reqwest::header::HeaderValue::from_str(value).map_err(|error| {
699 Error::Config(format!("invalid value for header {name}: {error}"))
701 })?;
702 map.insert(name, value);
703 }
704 Ok(map)
705}
706
707fn validate_base_url(base_url: &str) -> Result<(), Error> {
710 let invalid = || {
711 Error::Config(format!("invalid base URL {base_url:?}: expected an absolute http(s) URL with no query or fragment"))
712 };
713 let url = reqwest::Url::parse(base_url).map_err(|_| invalid())?;
714 let scheme_ok = matches!(url.scheme(), "http" | "https");
715 if !scheme_ok || url.host_str().is_none() || url.query().is_some() || url.fragment().is_some() {
716 return Err(invalid());
717 }
718 Ok(())
719}
720
721fn sdk_version_header() -> String {
722 format!("typesafe-client-rust/{}", env!("CARGO_PKG_VERSION"))
723}
724
725fn runtime_header() -> String {
726 format!(
727 "rust ({}; {})",
728 std::env::consts::OS,
729 std::env::consts::ARCH
730 )
731}