Skip to main content

elph_ai/api/
common.rs

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}