1use std::collections::HashMap;
2use std::future::Future;
3use std::pin::Pin;
4use std::sync::Arc;
5
6use anyhow::Result;
7use anyhow::anyhow;
8use reqwest::Client;
9use serde_json::Value;
10
11use crate::api::http_proxy::resolve_http_proxy_url_for_target;
12use crate::types::{AssistantMessage, AssistantMessageEvent, Model, OnPayloadCallback, OnResponseCallback};
13use crate::types::{ProviderEnv, ProviderResponse, StopReason, StreamOptions};
14use crate::utils::error_body::{error_body_from_response, format_provider_error, normalize_provider_error};
15use crate::utils::event_stream::AssistantMessageEventStream;
16use crate::utils::headers::{has_header, headers_to_record, merge_provider_headers};
17
18pub fn build_http_client(timeout_ms: Option<u64>) -> Result<Client> {
19 build_http_client_for_target(timeout_ms, None, None)
20}
21
22pub fn build_http_client_for_target(
23 timeout_ms: Option<u64>,
24 target_url: Option<&str>,
25 env: Option<&ProviderEnv>,
26) -> Result<Client> {
27 let mut builder = Client::builder();
28 if let Some(ms) = timeout_ms {
29 builder = builder.timeout(std::time::Duration::from_millis(ms));
30 }
31 if let Some(target_url) = target_url
32 && let Some(proxy_url) = resolve_http_proxy_url_for_target(target_url, env)?
33 {
34 let proxy = reqwest::Proxy::all(proxy_url.as_str())?;
35 builder = builder.proxy(proxy);
36 }
37 Ok(builder.build()?)
38}
39
40pub fn get_client_api_key(provider: &str, api_key: Option<&str>, headers: &HashMap<String, String>) -> Result<String> {
41 if let Some(key) = api_key {
42 return Ok(key.to_string());
43 }
44 if has_header(headers, "authorization") || has_header(headers, "cf-aig-authorization") {
45 return Ok("unused".to_string());
46 }
47 Err(anyhow!("No API key for provider: {provider}"))
48}
49
50pub async fn apply_on_payload(callback: Option<&OnPayloadCallback>, payload: Value, model: &Model) -> Value {
51 if let Some(cb) = callback {
52 let m = model.clone();
53 let original = payload.clone();
54 if let Some(next) = cb(payload, m).await {
55 return next;
56 }
57 return original;
58 }
59 payload
60}
61
62pub async fn apply_on_response(callback: Option<&OnResponseCallback>, response: ProviderResponse, model: &Model) {
63 if let Some(cb) = callback {
64 let m = model.clone();
65 cb(response, m).await;
66 }
67}
68
69pub fn merge_model_headers(model: &Model, options: Option<&StreamOptions>) -> HashMap<String, String> {
70 let base = model.headers.clone().unwrap_or_default();
71 merge_provider_headers(&base, options.and_then(|o| o.headers.as_ref()))
72}
73
74pub const REQUEST_ABORTED: &str = "Request aborted";
75
76pub fn is_request_aborted(token: &Option<tokio_util::sync::CancellationToken>) -> bool {
77 token.as_ref().is_some_and(|t| t.is_cancelled())
78}
79
80pub fn request_aborted_error() -> anyhow::Error {
81 anyhow!(REQUEST_ABORTED)
82}
83
84pub fn is_abort_error(error: &anyhow::Error) -> bool {
85 error.to_string() == REQUEST_ABORTED
86}
87
88pub fn with_trace_headers(request: reqwest::RequestBuilder) -> reqwest::RequestBuilder {
89 crate::trace::with_trace_headers(request)
90}
91
92pub async fn send_with_abort(
93 token: &Option<tokio_util::sync::CancellationToken>,
94 request: reqwest::RequestBuilder,
95) -> Result<reqwest::Response> {
96 if is_request_aborted(token) {
97 return Err(request_aborted_error());
98 }
99 let request = with_trace_headers(request);
100 match token {
101 Some(token) => {
102 let token = token.clone();
103 tokio::select! {
104 result = request.send() => result.map_err(Into::into),
105 _ = token.cancelled() => Err(request_aborted_error()),
106 }
107 }
108 None => request.send().await.map_err(Into::into),
109 }
110}
111
112pub fn finish_stream_error(
113 stream: &AssistantMessageEventStream,
114 output: &mut AssistantMessage,
115 error: anyhow::Error,
116 aborted: bool,
117) {
118 output.stop_reason = if aborted {
119 StopReason::Aborted
120 } else {
121 StopReason::Error
122 };
123 output.error_message = Some(format_provider_error(&normalize_provider_error(&error), None));
124 stream.push(AssistantMessageEvent::Error {
125 reason: output.stop_reason,
126 error: output.clone(),
127 });
128 stream.end();
129}
130
131pub async fn check_response_ok(response: reqwest::Response) -> Result<reqwest::Response> {
132 if response.status().is_success() {
133 return Ok(response);
134 }
135 let status = response.status();
136 let body = error_body_from_response(response).await;
137 Err(anyhow!("{status}: {body}"))
138}
139
140pub type StreamTask = Pin<Box<dyn Future<Output = ()> + Send>>;
141
142pub fn spawn_stream_task(fut: impl Future<Output = ()> + Send + 'static) -> StreamTask {
143 Box::pin(async move {
144 tokio::spawn(fut);
145 })
146}
147
148pub fn wrap_on_payload<F>(f: F) -> OnPayloadCallback
149where
150 F: Fn(Value, Model) -> Pin<Box<dyn Future<Output = Option<Value>> + Send>> + Send + Sync + 'static,
151{
152 Arc::new(f)
153}
154
155pub fn wrap_on_response<F>(f: F) -> OnResponseCallback
156where
157 F: Fn(ProviderResponse, Model) -> Pin<Box<dyn Future<Output = ()> + Send>> + Send + Sync + 'static,
158{
159 Arc::new(f)
160}
161
162pub async fn invoke_on_response_from_reqwest(
163 callback: Option<&OnResponseCallback>,
164 response: &reqwest::Response,
165 model: &Model,
166) {
167 let provider_response = ProviderResponse {
168 status: response.status().as_u16(),
169 headers: headers_to_record(response.headers()),
170 };
171 apply_on_response(callback, provider_response, model).await;
172}