diff --git a/Cargo.lock b/Cargo.lock index a18f9ce90..cdc90721d 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -1928,6 +1928,19 @@ dependencies = [ "zip", ] +[[package]] +name = "devolutions-gateway-ai" +version = "0.0.0" +dependencies = [ + "reqwest", + "secrecy", + "serde", + "serde_json", + "thiserror 2.0.20", + "tracing", + "url", +] + [[package]] name = "devolutions-gateway-generators" version = "0.0.0" @@ -7632,9 +7645,11 @@ dependencies = [ "agent-tunnel-proto", "anyhow", "assert_cmd", + "axum 0.8.9", "base64 0.23.1", "camino", "devolutions-gateway", + "devolutions-gateway-ai", "devolutions-gateway-task", "dynosaur", "escargot", @@ -7648,6 +7663,7 @@ dependencies = [ "network-scanner", "network-scanner-proto", "nonempty", + "parking_lot", "picky", "proxy-socks", "quinn", @@ -7670,6 +7686,7 @@ dependencies = [ "tokio-tungstenite", "tokio-util", "typed-builder", + "url", "uuid", "win-api-wrappers", "windows 0.61.3", diff --git a/crates/devolutions-gateway-ai/Cargo.toml b/crates/devolutions-gateway-ai/Cargo.toml new file mode 100644 index 000000000..9ce0c17dc --- /dev/null +++ b/crates/devolutions-gateway-ai/Cargo.toml @@ -0,0 +1,19 @@ +[package] +name = "devolutions-gateway-ai" +version = "0.0.0" +edition = "2024" +authors = ["Devolutions Inc. "] +description = "AI requests for Devolutions Gateway, one method per purpose" +publish = false + +[lints] +workspace = true + +[dependencies] +reqwest = { version = "0.12", default-features = false, features = ["json"] } +secrecy = "0.10" +serde = { version = "1", features = ["derive"] } +serde_json = "1" +thiserror = "2" +tracing = "0.1" +url = "2.5" diff --git a/crates/devolutions-gateway-ai/src/client.rs b/crates/devolutions-gateway-ai/src/client.rs new file mode 100644 index 000000000..d34bc3613 --- /dev/null +++ b/crates/devolutions-gateway-ai/src/client.rs @@ -0,0 +1,395 @@ +use std::fmt; +use std::time::Duration; + +use reqwest::header::HeaderValue; +use secrecy::{ExposeSecret as _, SecretString}; +use tracing::debug; +use url::Url; + +use crate::wire::{Api, anthropic, openai}; +use crate::{Error, Response, error}; + +/// Default of [`AiClientBuilder::request_timeout`]. +pub const DEFAULT_REQUEST_TIMEOUT: Duration = Duration::from_secs(10 * 60); + +/// AI provider behind an [`AiClient`]. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum Provider { + /// OpenAI chat completions; the default base URL is `https://api.openai.com/v1/`. + OpenAi, + /// Anthropic Messages; the default base URL is `https://api.anthropic.com/v1/`. + Anthropic, + /// Mistral chat completions; the default base URL is `https://api.mistral.ai/v1/`. + Mistral, + /// Google Gemini through its OpenAI-compatible endpoint; the default base URL is + /// `https://generativelanguage.googleapis.com/v1beta/openai/`. + Gemini, + /// Any other endpoint speaking OpenAI chat completions, such as a self-hosted model server; the base URL is required. + OpenAiCompatible, +} + +impl Provider { + /// Base URL used when the builder sets none; [`Provider::OpenAiCompatible`] has none. + pub fn default_base_url(self) -> Option { + let url = match self { + Self::OpenAi => "https://api.openai.com/v1/", + Self::Anthropic => "https://api.anthropic.com/v1/", + Self::Mistral => "https://api.mistral.ai/v1/", + Self::Gemini => "https://generativelanguage.googleapis.com/v1beta/openai/", + Self::OpenAiCompatible => return None, + }; + + Some(Url::parse(url).expect("default base URLs are valid")) + } + + fn api(self) -> Api { + match self { + // OpenAI's newer models only accept `max_completion_tokens`; other servers only know `max_tokens`. + Self::OpenAi => Api::OpenAiChat(openai::TokenLimit::MaxCompletionTokens), + Self::Mistral | Self::Gemini | Self::OpenAiCompatible => Api::OpenAiChat(openai::TokenLimit::MaxTokens), + Self::Anthropic => Api::AnthropicMessages, + } + } +} + +/// Error returned by [`AiClientBuilder::build`]. +#[derive(Debug, thiserror::Error)] +#[non_exhaustive] +pub enum BuildError { + #[error("AI provider is missing")] + MissingProvider, + #[error("AI model is missing")] + MissingModel, + #[error("API key is missing for AI provider {0:?}")] + MissingApiKey(Provider), + /// The API key cannot be sent in an HTTP header, such as a key holding a line break. + #[error("API key is not a valid HTTP header value")] + InvalidApiKey, + #[error("base URL is missing for AI provider {0:?}")] + MissingBaseUrl(Provider), + #[error("base URL scheme must be http or https")] + UnsupportedBaseUrl, + #[error("HTTP client is missing")] + MissingHttpClient, +} + +/// Builds an [`AiClient`]; every setting is checked by [`AiClientBuilder::build`]. +#[derive(Debug, Default)] +pub struct AiClientBuilder { + provider: Option, + model: Option, + api_key: Option, + base_url: Option, + http_client: Option, + request_timeout: Option, +} + +impl AiClientBuilder { + #[must_use] + pub fn provider(mut self, provider: Provider) -> Self { + self.provider = Some(provider); + self + } + + /// Model identifier, passed to the provider as is. + #[must_use] + pub fn model(mut self, model: impl Into) -> Self { + self.model = Some(model.into()); + self + } + + #[must_use] + pub fn api_key(mut self, api_key: impl Into) -> Self { + self.api_key = Some(api_key.into()); + self + } + + /// Overrides [`Provider::default_base_url`]; required for [`Provider::OpenAiCompatible`]. + #[must_use] + pub fn base_url(mut self, base_url: Url) -> Self { + self.base_url = Some(base_url); + self + } + + /// Client used for every request, so the caller's proxy and TLS policy apply. + #[must_use] + pub fn http_client(mut self, http_client: reqwest::Client) -> Self { + self.http_client = Some(http_client); + self + } + + /// Longest time one request may take, answer included; the default is [`DEFAULT_REQUEST_TIMEOUT`]. + /// + /// Requests are not streamed, so it must leave the model time to write its whole answer. + #[must_use] + pub fn request_timeout(mut self, request_timeout: Duration) -> Self { + self.request_timeout = Some(request_timeout); + self + } + + /// Checks every setting, so that no request of the client can fail because of them. + pub fn build(self) -> Result { + let provider = self.provider.ok_or(BuildError::MissingProvider)?; + + let model = self + .model + .filter(|model| !model.trim().is_empty()) + .ok_or(BuildError::MissingModel)?; + + let api_key = self + .api_key + .filter(|api_key| !api_key.expose_secret().is_empty()) + .ok_or(BuildError::MissingApiKey(provider))?; + + if HeaderValue::from_str(api_key.expose_secret()).is_err() { + return Err(BuildError::InvalidApiKey); + } + + let base_url = resolve_base_url(provider, self.base_url)?; + let http_client = self.http_client.ok_or(BuildError::MissingHttpClient)?; + + Ok(AiClient { + provider, + model, + base_url, + api_key, + http_client, + request_timeout: self.request_timeout.unwrap_or(DEFAULT_REQUEST_TIMEOUT), + }) + } +} + +fn resolve_base_url(provider: Provider, base_url: Option) -> Result { + let base_url = base_url + .or_else(|| provider.default_base_url()) + .ok_or(BuildError::MissingBaseUrl(provider))?; + + if !matches!(base_url.scheme(), "http" | "https") { + return Err(BuildError::UnsupportedBaseUrl); + } + + Ok(base_url) +} + +/// Sends the requests of every purpose to one AI provider. +#[derive(Clone)] +pub struct AiClient { + provider: Provider, + model: String, + base_url: Url, + api_key: SecretString, + http_client: reqwest::Client, + request_timeout: Duration, +} + +impl fmt::Debug for AiClient { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct("AiClient") + .field("provider", &self.provider) + .field("model", &self.model) + .field("base_url", &self.base_url.as_str()) + .field("request_timeout", &self.request_timeout) + .finish_non_exhaustive() + } +} + +/// Text completion request; a purpose is defined by its prompt and the parser of the answer. +pub(crate) struct Prompt<'a> { + /// Instructions of the purpose. + pub(crate) system: &'a str, + /// Data the purpose works on, such as a session transcript. + pub(crate) input: &'a str, + pub(crate) max_output_tokens: u32, +} + +impl AiClient { + pub fn builder() -> AiClientBuilder { + AiClientBuilder::default() + } + + /// Sends one completion request and returns the text of the answer. + /// + /// An answer cut at the output token limit is [`Error::Truncated`], because no purpose can use a partial answer. + /// The `` blocks some models write before their answer are removed. + pub(crate) async fn complete(&self, prompt: &Prompt<'_>) -> Result, Error> { + debug!( + provider = ?self.provider, + model = %self.model, + base_url = %self.base_url, + input_len = prompt.input.len(), + max_output_tokens = prompt.max_output_tokens, + "Send AI completion request" + ); + + let api = self.provider.api(); + let api_key = self.api_key.expose_secret(); + + let request = match api { + Api::OpenAiChat(limit) => { + openai::request(&self.http_client, &self.base_url, api_key, &self.model, prompt, limit) + } + Api::AnthropicMessages => { + anthropic::request(&self.http_client, &self.base_url, api_key, &self.model, prompt) + } + }; + + let response = request + .timeout(self.request_timeout) + .send() + .await + .map_err(|error| error::transport(&error, api_key))?; + + let status = response.status(); + let body = response + .bytes() + .await + .map_err(|error| error::transport(&error, api_key))?; + + if !status.is_success() { + return Err(error::status(status, &body, api_key)); + } + + let completion = match api { + Api::OpenAiChat(_) => openai::parse(&body)?, + Api::AnthropicMessages => anthropic::parse(&body)?, + }; + + let text = strip_think_blocks(completion.text); + + debug!( + output_len = text.len(), + truncated = completion.truncated, + model = ?completion.model, + usage = ?completion.usage, + "Received AI completion" + ); + + if completion.truncated { + return Err(Error::Truncated { + usage: completion.usage, + }); + } + + Ok(Response { + output: text, + model: completion.model, + usage: completion.usage, + }) + } +} + +/// Removes the `…` blocks that some models write before their answer; an unclosed block runs to the end. +fn strip_think_blocks(text: String) -> String { + const OPEN: &str = ""; + const CLOSE: &str = ""; + + if find_ignore_ascii_case(&text, OPEN).is_none() { + return text; + } + + let mut stripped = String::with_capacity(text.len()); + let mut rest = text.as_str(); + + while let Some(open) = find_ignore_ascii_case(rest, OPEN) { + stripped.push_str(&rest[..open]); + + let inside = &rest[open + OPEN.len()..]; + let Some(close) = find_ignore_ascii_case(inside, CLOSE) else { + return stripped; + }; + + rest = &inside[close + CLOSE.len()..]; + } + + stripped.push_str(rest); + stripped +} + +// The needle is ASCII, so a match always starts on a character boundary. +fn find_ignore_ascii_case(haystack: &str, needle: &str) -> Option { + haystack + .as_bytes() + .windows(needle.len()) + .position(|window| window.eq_ignore_ascii_case(needle.as_bytes())) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn builder_debug_hides_the_api_key() { + let builder = AiClient::builder() + .provider(Provider::OpenAi) + .model("gpt-test") + .api_key("sk-very-secret"); + + assert!(!format!("{builder:?}").contains("sk-very-secret")); + } + + #[test] + fn default_base_urls() { + for (provider, expected) in [ + (Provider::OpenAi, "https://api.openai.com/v1/"), + (Provider::Anthropic, "https://api.anthropic.com/v1/"), + (Provider::Mistral, "https://api.mistral.ai/v1/"), + ( + Provider::Gemini, + "https://generativelanguage.googleapis.com/v1beta/openai/", + ), + ] { + let resolved = resolve_base_url(provider, None).expect("default base URL"); + assert_eq!(resolved.as_str(), expected); + } + + assert!(matches!( + resolve_base_url(Provider::OpenAiCompatible, None), + Err(BuildError::MissingBaseUrl(Provider::OpenAiCompatible)) + )); + } + + #[test] + fn think_blocks_are_removed() { + for (text, expected) in [ + ("answer", "answer"), + ("plananswer", "answer"), + ("\nplan\n\nanswer\n", "\nanswer\n"), + ("a1b2c", "abc"), + ("éplan", "é"), + ("answer", "answer"), + ] { + assert_eq!(strip_think_blocks(text.to_owned()), expected, "{text:?}"); + } + } + + #[test] + fn base_url_overrides_default() { + let custom = Url::parse("https://proxy.example/anthropic/").expect("valid URL"); + + let resolved = resolve_base_url(Provider::Anthropic, Some(custom.clone())).expect("custom base URL"); + + assert_eq!(resolved, custom); + } + + #[test] + fn base_url_must_be_http() { + let url = Url::parse("ftp://files.example/v1/").expect("valid URL"); + + assert!(matches!( + resolve_base_url(Provider::OpenAiCompatible, Some(url)), + Err(BuildError::UnsupportedBaseUrl) + )); + } + + #[test] + fn api_key_must_fit_in_a_header() { + let result = AiClient::builder() + .provider(Provider::OpenAi) + .model("gpt-test") + .api_key("sk-line\nbreak") + .http_client(reqwest::Client::new()) + .build(); + + assert!(matches!(result, Err(BuildError::InvalidApiKey)), "{result:?}"); + } +} diff --git a/crates/devolutions-gateway-ai/src/error.rs b/crates/devolutions-gateway-ai/src/error.rs new file mode 100644 index 000000000..82f625366 --- /dev/null +++ b/crates/devolutions-gateway-ai/src/error.rs @@ -0,0 +1,146 @@ +use reqwest::StatusCode; + +use crate::{Usage, wire}; + +/// Error returned when a purpose request fails; no variant ever holds the API key. +#[derive(Debug, thiserror::Error)] +#[non_exhaustive] +pub enum Error { + /// The request got no answer, such as on a connection failure or a timeout. + /// + /// The underlying error is not kept as `source()`, because its chain could expose the key unredacted. + #[error("AI provider request failed: {message}")] + Transport { message: String }, + /// The provider answered with an error status. + #[error("AI provider answered HTTP {status}: {message}")] + Status { status: u16, message: String }, + /// The provider answered with a body that is not in the format of its API. + #[error("AI provider answer is not valid: {reason}")] + InvalidResponse { reason: String }, + /// The answer reached the output token limit, so its end is missing: send a shorter input instead. + /// + /// The provider still counts the tokens of the request in `usage`. + #[error("AI answer was cut at the output token limit")] + Truncated { usage: Option }, + /// The answer does not follow the output format the purpose asked for. + #[error("AI answer is not in the expected format: {reason}")] + InvalidOutput { reason: String }, +} + +impl Error { + /// Returns `true` when sending the same request again later may succeed, such as after a rate limit. + pub fn is_transient(&self) -> bool { + match self { + Self::Transport { .. } => true, + // Request timeout, rate limit, and server errors, including Anthropic's 529 "overloaded". + Self::Status { status, .. } => matches!(*status, 408 | 429 | 500..), + Self::InvalidResponse { .. } | Self::Truncated { .. } | Self::InvalidOutput { .. } => false, + } + } +} + +pub(crate) fn transport(error: &dyn std::error::Error, api_key: &str) -> Error { + let mut message = error.to_string(); + let mut source = error.source(); + while let Some(cause) = source { + message.push_str(": "); + message.push_str(&cause.to_string()); + source = cause.source(); + } + + Error::Transport { + message: redact(message, api_key), + } +} + +pub(crate) fn status(status: StatusCode, body: &[u8], api_key: &str) -> Error { + let message = + wire::error_message(body).unwrap_or_else(|| status.canonical_reason().unwrap_or("unknown status").to_owned()); + + Error::Status { + status: status.as_u16(), + message: redact(message, api_key), + } +} + +// Provider error bodies may echo the API key back. +fn redact(message: String, api_key: &str) -> String { + if api_key.is_empty() { + message + } else { + message.replace(api_key, "[REDACTED]") + } +} + +#[cfg(test)] +mod tests { + use super::*; + + const API_KEY: &str = "sk-very-secret"; + + #[test] + fn transport_error_redacts_api_key() { + let error = std::io::Error::other(format!("invalid key {API_KEY} provided")); + + let error = transport(&error, API_KEY); + + assert!(!error.to_string().contains(API_KEY)); + assert!(!format!("{error:?}").contains(API_KEY)); + assert!(error.to_string().contains("[REDACTED]")); + } + + #[test] + fn status_error_redacts_api_key_echoed_by_the_provider() { + let body = format!(r#"{{"error":{{"message":"Incorrect API key provided: {API_KEY}"}}}}"#); + + let error = status(StatusCode::UNAUTHORIZED, body.as_bytes(), API_KEY); + + assert!(matches!(error, Error::Status { status: 401, .. }), "{error:?}"); + assert!(!format!("{error:?}").contains(API_KEY)); + assert!(error.to_string().contains("[REDACTED]")); + } + + #[test] + fn status_error_without_provider_message_uses_the_reason() { + let error = status(StatusCode::BAD_GATEWAY, b"proxy error", API_KEY); + + assert_eq!(error.to_string(), "AI provider answered HTTP 502: Bad Gateway"); + } + + #[test] + fn transient_errors() { + let status = |status| Error::Status { + status, + message: "failed".to_owned(), + }; + + for error in [ + Error::Transport { + message: "connection refused".to_owned(), + }, + status(408), + status(429), + status(500), + status(503), + status(529), + ] { + assert!(error.is_transient(), "{error:?}"); + } + + for error in [ + status(400), + status(401), + status(403), + status(404), + Error::InvalidResponse { + reason: "syntax error".to_owned(), + }, + Error::Truncated { usage: None }, + Error::InvalidOutput { + reason: "no valid line".to_owned(), + }, + ] { + assert!(!error.is_transient(), "{error:?}"); + } + } +} diff --git a/crates/devolutions-gateway-ai/src/lib.rs b/crates/devolutions-gateway-ai/src/lib.rs new file mode 100644 index 000000000..3dddfae68 --- /dev/null +++ b/crates/devolutions-gateway-ai/src/lib.rs @@ -0,0 +1,28 @@ +//! AI requests for Devolutions Gateway, one method per purpose. +//! +//! A purpose is one job Gateway gives to an AI model, such as listing the actions of a session transcript with +//! [`AiClient::describe_session_actions`]. +//! Each purpose owns its prompt, the version of that prompt, and the parser of the answer, so Gateway gets typed +//! results and never writes a prompt or reads raw model text. +//! Every purpose goes through one [`AiClient`], which holds the provider settings, and returns a [`Response`], which +//! also tells which model answered and how many tokens the request used. +//! +//! Each provider is reached through its own HTTP API: OpenAI chat completions (also spoken by Mistral, Gemini and many +//! others) or Anthropic Messages. Only the few fields a single text completion needs are modeled. +//! Requests are not streamed and are bounded by a timeout. +//! +//! A new purpose is a module like [`session_actions`]: a prompt and its `PROMPT_VERSION`, a request builder returned by +//! a new [`AiClient`] method, and a parser turning the answer into typed output. + +mod client; +mod error; +mod response; +pub mod session_actions; +mod wire; + +pub use reqwest; +pub use secrecy; + +pub use self::client::{AiClient, AiClientBuilder, BuildError, DEFAULT_REQUEST_TIMEOUT, Provider}; +pub use self::error::Error; +pub use self::response::{Response, Usage}; diff --git a/crates/devolutions-gateway-ai/src/response.rs b/crates/devolutions-gateway-ai/src/response.rs new file mode 100644 index 000000000..681db2d2c --- /dev/null +++ b/crates/devolutions-gateway-ai/src/response.rs @@ -0,0 +1,46 @@ +use core::ops::Add; + +use crate::Error; + +/// Answer to a purpose request. +#[derive(Debug, Clone, PartialEq)] +pub struct Response { + /// Result of the purpose, such as the actions of a session. + pub output: T, + /// Model that answered, as reported by the provider. + /// + /// It is often more precise than the requested model, such as a dated version of it. + pub model: Option, + /// Tokens the provider counted for the request, when it reports them. + pub usage: Option, +} + +impl Response { + pub(crate) fn try_map(self, f: impl FnOnce(T) -> Result) -> Result, Error> { + Ok(Response { + output: f(self.output)?, + model: self.model, + usage: self.usage, + }) + } +} + +/// Tokens counted by the provider for one or more requests; providers bill by them. +#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)] +pub struct Usage { + /// Tokens of the prompt and the input. + pub input_tokens: u64, + /// Tokens of the answer. + pub output_tokens: u64, +} + +impl Add for Usage { + type Output = Self; + + fn add(self, other: Self) -> Self { + Self { + input_tokens: self.input_tokens.saturating_add(other.input_tokens), + output_tokens: self.output_tokens.saturating_add(other.output_tokens), + } + } +} diff --git a/crates/devolutions-gateway-ai/src/session_actions.rs b/crates/devolutions-gateway-ai/src/session_actions.rs new file mode 100644 index 000000000..17f4d3b27 --- /dev/null +++ b/crates/devolutions-gateway-ai/src/session_actions.rs @@ -0,0 +1,273 @@ +//! Purpose: list what the user did in a session transcript, with [`AiClient::describe_session_actions`]. + +use std::collections::BTreeMap; +use std::fmt; +use std::time::Duration; + +use serde::Deserialize; +use tracing::warn; + +use crate::client::{AiClient, Prompt}; +use crate::{Error, Response}; + +/// Version of the prompt of this purpose. +/// +/// Bump it whenever the prompt changes, so readers of the results know which prompt produced them. +pub const PROMPT_VERSION: &str = "session-actions-1"; + +const PROMPT: &str = r#"You read the transcript of a remote session and list what the user did. + +Input: each transcript line starts with the elapsed time since the session started, in seconds, between square brackets. Example: `[12.5] ls -la`. + +Output: JSON Lines only. Write one JSON object per line and nothing else: no prose, no Markdown, no code fences. +Each object has these fields: +- "offsetSeconds": number. Elapsed seconds when the action started, taken from the transcript. +- "description": string. A short past-tense sentence naming the action, like "Listed directory contents". +- "object": string, optional. The main thing acted on, like a file path, host, service, or account. +- "parameters": object, optional. Every value is a string. Important details, like the exact command. + +Example output line: +{"offsetSeconds":12.5,"description":"Listed directory contents","object":"/var/log","parameters":{"Command":"ls -la /var/log"}} + +Rules: +- Write one line per meaningful user action, in time order. Merge the keystrokes of one command into one action. +- Ignore noise, such as prompt redraws, cursor movement, and output that has no user action. +- Never copy passwords, secrets, or tokens. Write "[redacted]" instead. +- If the user did nothing, write nothing."#; + +/// The same limit DVLS and RDM use for Claude. Reasoning models count their reasoning in it, so it cannot be small. +const DEFAULT_MAX_OUTPUT_TOKENS: u32 = 16_000; + +/// One user action found in a session transcript. +#[derive(Debug, Clone, PartialEq)] +pub struct Action { + /// Elapsed time since the start of the session, as reported by the model. + pub offset: Duration, + /// Short past-tense sentence naming the action, never empty. + pub description: String, + /// Main thing acted on, such as a file path, host, service, or account; never an empty string. + pub object: Option, + /// Important details, such as the exact command, keyed by name. + pub parameters: BTreeMap, +} + +impl AiClient { + /// Asks the model which actions the user performed in a session transcript. + /// + /// Each line of `transcript` must start with the elapsed time in seconds between square brackets, such as + /// `[12.5] ls`. + pub fn describe_session_actions<'a>(&'a self, transcript: &'a str) -> DescribeSessionActions<'a> { + DescribeSessionActions { + client: self, + transcript, + max_output_tokens: DEFAULT_MAX_OUTPUT_TOKENS, + } + } +} + +/// Request built by [`AiClient::describe_session_actions`]. +#[must_use = "the request is sent only by `send`"] +pub struct DescribeSessionActions<'a> { + client: &'a AiClient, + transcript: &'a str, + max_output_tokens: u32, +} + +// The transcript holds session data, so only its length is printed. +impl fmt::Debug for DescribeSessionActions<'_> { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct("DescribeSessionActions") + .field("client", self.client) + .field("transcript_len", &self.transcript.len()) + .field("max_output_tokens", &self.max_output_tokens) + .finish() + } +} + +impl DescribeSessionActions<'_> { + /// Upper bound of tokens in the answer; the default is 16000. + pub fn max_output_tokens(mut self, max_output_tokens: u32) -> Self { + self.max_output_tokens = max_output_tokens; + self + } + + /// Returns the actions in the order of the answer. + /// + /// Invalid lines in the answer are skipped with a warning. + /// The answer is [`Error::InvalidOutput`] only when it has lines but none of them is a valid action. + /// An answer cut at the output token limit is [`Error::Truncated`]: send a shorter transcript instead. + pub async fn send(self) -> Result>, Error> { + let prompt = Prompt { + system: PROMPT, + input: self.transcript, + max_output_tokens: self.max_output_tokens, + }; + + self.client + .complete(&prompt) + .await? + .try_map(|answer| parse_actions(&answer)) + } +} + +#[derive(Deserialize)] +#[serde(rename_all = "camelCase")] +struct ActionLine { + offset_seconds: f64, + description: String, + #[serde(default)] + object: Option, + #[serde(default)] + parameters: BTreeMap, +} + +fn parse_actions(answer: &str) -> Result, Error> { + let mut actions = Vec::new(); + let mut invalid_lines = 0usize; + + for (index, line) in answer.lines().enumerate() { + let line = line.trim(); + + if line.is_empty() || line.starts_with("```") { + continue; + } + + match parse_action_line(line) { + Ok(action) => actions.push(action), + Err(reason) => { + invalid_lines += 1; + warn!(line_number = index + 1, %reason, "Skipped invalid AI action line"); + } + } + } + + if actions.is_empty() && invalid_lines > 0 { + return Err(Error::InvalidOutput { + reason: format!("no valid action line, {invalid_lines} invalid lines"), + }); + } + + Ok(actions) +} + +// The reason never quotes the line, because the line may contain session data. +fn parse_action_line(line: &str) -> Result { + let parsed: ActionLine = serde_json::from_str(line) + .map_err(|error| format!("{:?} error at column {}", error.classify(), error.column()))?; + + let offset = Duration::try_from_secs_f64(parsed.offset_seconds).map_err(|_| "invalid offsetSeconds".to_owned())?; + + let description = parsed.description.trim(); + if description.is_empty() { + return Err("empty description".to_owned()); + } + + Ok(Action { + offset, + description: description.to_owned(), + object: parsed + .object + .map(|object| object.trim().to_owned()) + .filter(|object| !object.is_empty()), + parameters: parsed.parameters, + }) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn parses_valid_lines() { + let answer = concat!( + "{\"offsetSeconds\":1.5,\"description\":\"Listed files\",\"object\":\"/var/log\",\"parameters\":{\"Command\":\"ls\"}}\n", + "\n", + "{\"offsetSeconds\":3,\"description\":\"Opened a shell\"}\n", + ); + + let actions = parse_actions(answer).expect("valid answer"); + + assert_eq!( + actions, + vec![ + Action { + offset: Duration::from_millis(1500), + description: "Listed files".to_owned(), + object: Some("/var/log".to_owned()), + parameters: BTreeMap::from([("Command".to_owned(), "ls".to_owned())]), + }, + Action { + offset: Duration::from_secs(3), + description: "Opened a shell".to_owned(), + object: None, + parameters: BTreeMap::new(), + }, + ] + ); + } + + #[test] + fn skips_code_fences_and_invalid_lines() { + let answer = concat!( + "```jsonl\n", + "{\"offsetSeconds\":1,\"description\":\"Listed files\"}\n", + "Here are the actions:\n", + "{\"offsetSeconds\":-1,\"description\":\"Negative offset\"}\n", + "{\"offsetSeconds\":2,\"description\":\" \"}\n", + "{\"offsetSeconds\":2,\"description\":\"Numeric parameter\",\"parameters\":{\"Count\":5}}\n", + "```\n", + ); + + let actions = parse_actions(answer).expect("one valid line"); + + assert_eq!(actions.len(), 1); + assert_eq!(actions[0].description, "Listed files"); + } + + #[test] + fn empty_answer_means_no_action() { + assert_eq!(parse_actions("").expect("empty answer"), Vec::new()); + assert_eq!(parse_actions("\n```\n```\n").expect("only fences"), Vec::new()); + } + + #[test] + fn answer_without_valid_line_is_invalid_output() { + let error = parse_actions("not json\n{\"description\":\"no offset\"}\n").expect_err("no valid line"); + + assert!(matches!(error, Error::InvalidOutput { .. }), "{error:?}"); + assert!(error.to_string().contains("2 invalid lines"), "{error}"); + } + + #[test] + fn invalid_line_reason_does_not_quote_the_line() { + let reason = + parse_action_line("{\"offsetSeconds\":1,\"description\":\"secret-value\",\"parameters\":{\"a\":1}}") + .expect_err("numeric parameter"); + + assert!(!reason.contains("secret-value")); + } + + #[test] + fn prompt_asks_for_the_parsed_fields() { + for field in ["offsetSeconds", "description", "object", "parameters", "JSON Lines"] { + assert!(PROMPT.contains(field), "prompt is missing {field}"); + } + } + + #[test] + fn request_debug_hides_the_transcript() { + let client = AiClient::builder() + .provider(crate::Provider::OpenAi) + .model("gpt-test") + .api_key("sk-very-secret") + .http_client(reqwest::Client::new()) + .build() + .expect("valid settings"); + + let debug = format!("{:?}", client.describe_session_actions("[1] secret-command")); + + assert!(!debug.contains("secret-command")); + assert!(!debug.contains("sk-very-secret")); + assert!(debug.contains("transcript_len: 18")); + } +} diff --git a/crates/devolutions-gateway-ai/src/wire/anthropic.rs b/crates/devolutions-gateway-ai/src/wire/anthropic.rs new file mode 100644 index 000000000..42f34a1ef --- /dev/null +++ b/crates/devolutions-gateway-ai/src/wire/anthropic.rs @@ -0,0 +1,121 @@ +//! Anthropic Messages. + +use serde::{Deserialize, Serialize}; +use url::Url; + +use super::{Completion, Message, endpoint, parse_body, usage}; +use crate::Error; +use crate::client::Prompt; + +const VERSION: &str = "2023-06-01"; + +pub(crate) fn request( + http_client: &reqwest::Client, + base_url: &Url, + api_key: &str, + model: &str, + prompt: &Prompt<'_>, +) -> reqwest::RequestBuilder { + http_client + .post(endpoint(base_url, "messages")) + .header("x-api-key", api_key) + .header("anthropic-version", VERSION) + .json(&MessagesRequest { + model, + system: prompt.system, + messages: [Message { + role: "user", + content: prompt.input, + }], + max_tokens: prompt.max_output_tokens, + }) +} + +pub(crate) fn parse(body: &[u8]) -> Result { + let response: MessagesResponse = parse_body(body)?; + + let text = response + .content + .into_iter() + .filter_map(|block| match block { + ContentBlock::Text { text } => Some(text), + ContentBlock::Other => None, + }) + .collect::>() + .join("\n"); + + Ok(Completion { + text, + truncated: response.stop_reason.as_deref() == Some("max_tokens"), + model: response.model, + usage: response + .usage + .and_then(|counts| usage(counts.input_tokens, counts.output_tokens)), + }) +} + +#[derive(Serialize)] +struct MessagesRequest<'a> { + model: &'a str, + system: &'a str, + messages: [Message<'a>; 1], + max_tokens: u32, +} + +#[derive(Deserialize)] +struct MessagesResponse { + model: Option, + content: Vec, + stop_reason: Option, + usage: Option, +} + +#[derive(Deserialize)] +#[serde(tag = "type", rename_all = "snake_case")] +enum ContentBlock { + Text { + text: String, + }, + #[serde(other)] + Other, +} + +#[derive(Deserialize)] +struct MessagesUsage { + input_tokens: Option, + output_tokens: Option, +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::Usage; + + #[test] + fn parses_text_blocks_model_and_usage() { + let completion = parse( + br#"{"model":"claude-test-20260101","content":[{"type":"text","text":"a"},{"type":"thinking","thinking":"hidden"},{"type":"text","text":"b"}],"stop_reason":"end_turn","usage":{"input_tokens":10,"output_tokens":20}}"#, + ) + .expect("valid answer"); + + assert_eq!(completion.text, "a\nb"); + assert!(!completion.truncated); + assert_eq!(completion.model.as_deref(), Some("claude-test-20260101")); + assert_eq!( + completion.usage, + Some(Usage { + input_tokens: 10, + output_tokens: 20 + }) + ); + } + + #[test] + fn max_tokens_stop_reason_is_truncated() { + let completion = parse(br#"{"content":[{"type":"text","text":"partial"}],"stop_reason":"max_tokens"}"#) + .expect("valid answer"); + + assert!(completion.truncated); + assert_eq!(completion.usage, None); + } +} diff --git a/crates/devolutions-gateway-ai/src/wire/mod.rs b/crates/devolutions-gateway-ai/src/wire/mod.rs new file mode 100644 index 000000000..d20cd7128 --- /dev/null +++ b/crates/devolutions-gateway-ai/src/wire/mod.rs @@ -0,0 +1,109 @@ +//! Request and answer formats of the provider HTTP APIs, reduced to what a single text completion needs. + +pub(crate) mod anthropic; +pub(crate) mod openai; + +use serde::de::DeserializeOwned; +use serde::{Deserialize, Serialize}; +use url::Url; + +use crate::{Error, Usage}; + +/// HTTP API spoken by a provider. +#[derive(Debug, Clone, Copy)] +pub(crate) enum Api { + OpenAiChat(openai::TokenLimit), + AnthropicMessages, +} + +/// Answer of a provider, in the terms shared by every API. +pub(crate) struct Completion { + pub(crate) text: String, + /// The answer reached the output token limit. + pub(crate) truncated: bool, + pub(crate) model: Option, + pub(crate) usage: Option, +} + +#[derive(Serialize)] +struct Message<'a> { + role: &'static str, + content: &'a str, +} + +/// Error body shared by the OpenAI-style and Anthropic APIs. +#[derive(Deserialize)] +struct ErrorBody { + error: ErrorDetail, +} + +#[derive(Deserialize)] +struct ErrorDetail { + message: String, +} + +pub(crate) fn error_message(body: &[u8]) -> Option { + serde_json::from_slice::(body) + .ok() + .map(|body| body.error.message) +} + +fn endpoint(base_url: &Url, path: &str) -> String { + format!("{}/{path}", base_url.as_str().trim_end_matches('/')) +} + +// The reason never quotes the body, because the answer may contain session data. +fn parse_body(body: &[u8]) -> Result { + serde_json::from_slice(body).map_err(|error| Error::InvalidResponse { + reason: format!( + "{:?} error at line {} column {}", + error.classify(), + error.line(), + error.column() + ), + }) +} + +fn usage(input_tokens: Option, output_tokens: Option) -> Option { + Some(Usage { + input_tokens: input_tokens?, + output_tokens: output_tokens?, + }) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn endpoint_joins_with_or_without_trailing_slash() { + for base_url in ["https://api.example/v1/", "https://api.example/v1"] { + let base_url = Url::parse(base_url).expect("valid URL"); + assert_eq!( + endpoint(&base_url, "chat/completions"), + "https://api.example/v1/chat/completions" + ); + } + } + + #[test] + fn invalid_response_reason_does_not_quote_the_body() { + let error = parse_body::>(br#"["secret-value"]"#).expect_err("invalid body"); + + assert!(matches!(error, Error::InvalidResponse { .. })); + assert!(!error.to_string().contains("secret-value")); + } + + #[test] + fn usage_needs_both_counts() { + assert_eq!( + usage(Some(1), Some(2)), + Some(Usage { + input_tokens: 1, + output_tokens: 2 + }) + ); + assert_eq!(usage(Some(1), None), None); + assert_eq!(usage(None, Some(2)), None); + } +} diff --git a/crates/devolutions-gateway-ai/src/wire/openai.rs b/crates/devolutions-gateway-ai/src/wire/openai.rs new file mode 100644 index 000000000..a59825882 --- /dev/null +++ b/crates/devolutions-gateway-ai/src/wire/openai.rs @@ -0,0 +1,143 @@ +//! OpenAI chat completions, also spoken by Mistral, Gemini, and many other providers. + +use serde::{Deserialize, Serialize}; +use url::Url; + +use super::{Completion, Message, endpoint, parse_body, usage}; +use crate::Error; +use crate::client::Prompt; + +/// Field holding the output token limit in the request. +#[derive(Debug, Clone, Copy)] +pub(crate) enum TokenLimit { + /// `max_completion_tokens`, the only one the newer OpenAI models accept. + MaxCompletionTokens, + /// `max_tokens`, the only one most other servers know. + MaxTokens, +} + +pub(crate) fn request( + http_client: &reqwest::Client, + base_url: &Url, + api_key: &str, + model: &str, + prompt: &Prompt<'_>, + limit: TokenLimit, +) -> reqwest::RequestBuilder { + let (max_completion_tokens, max_tokens) = match limit { + TokenLimit::MaxCompletionTokens => (Some(prompt.max_output_tokens), None), + TokenLimit::MaxTokens => (None, Some(prompt.max_output_tokens)), + }; + + http_client + .post(endpoint(base_url, "chat/completions")) + .bearer_auth(api_key) + .json(&ChatRequest { + model, + messages: [ + Message { + role: "system", + content: prompt.system, + }, + Message { + role: "user", + content: prompt.input, + }, + ], + max_completion_tokens, + max_tokens, + }) +} + +pub(crate) fn parse(body: &[u8]) -> Result { + let response: ChatResponse = parse_body(body)?; + let choice = response.choices.into_iter().next(); + + Ok(Completion { + truncated: choice + .as_ref() + .is_some_and(|choice| choice.finish_reason.as_deref() == Some("length")), + text: choice.and_then(|choice| choice.message.content).unwrap_or_default(), + model: response.model, + usage: response + .usage + .and_then(|counts| usage(counts.prompt_tokens, counts.completion_tokens)), + }) +} + +#[derive(Serialize)] +struct ChatRequest<'a> { + model: &'a str, + messages: [Message<'a>; 2], + #[serde(skip_serializing_if = "Option::is_none")] + max_completion_tokens: Option, + #[serde(skip_serializing_if = "Option::is_none")] + max_tokens: Option, +} + +#[derive(Deserialize)] +struct ChatResponse { + model: Option, + choices: Vec, + usage: Option, +} + +#[derive(Deserialize)] +struct ChatChoice { + message: ChatAnswer, + finish_reason: Option, +} + +#[derive(Deserialize)] +struct ChatAnswer { + content: Option, +} + +#[derive(Deserialize)] +struct ChatUsage { + prompt_tokens: Option, + completion_tokens: Option, +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::Usage; + + #[test] + fn parses_text_model_and_usage() { + let completion = parse( + br#"{"model":"gpt-4o-2024-08-06","choices":[{"message":{"role":"assistant","content":"hello"},"finish_reason":"stop"}],"usage":{"prompt_tokens":10,"completion_tokens":20,"total_tokens":30}}"#, + ) + .expect("valid answer"); + + assert_eq!(completion.text, "hello"); + assert!(!completion.truncated); + assert_eq!(completion.model.as_deref(), Some("gpt-4o-2024-08-06")); + assert_eq!( + completion.usage, + Some(Usage { + input_tokens: 10, + output_tokens: 20 + }) + ); + } + + #[test] + fn length_finish_reason_is_truncated() { + let completion = parse(br#"{"choices":[{"message":{"content":"partial"},"finish_reason":"length"}]}"#) + .expect("valid answer"); + + assert!(completion.truncated); + } + + #[test] + fn model_usage_and_content_are_optional() { + let completion = parse(br#"{"choices":[{"message":{"content":null}}]}"#).expect("valid answer"); + + assert_eq!(completion.text, ""); + assert!(!completion.truncated); + assert_eq!(completion.model, None); + assert_eq!(completion.usage, None); + } +} diff --git a/testsuite/Cargo.toml b/testsuite/Cargo.toml index a40691073..fd029c343 100644 --- a/testsuite/Cargo.toml +++ b/testsuite/Cargo.toml @@ -35,8 +35,10 @@ agent-sysevent-codes.path = "../crates/agent-sysevent-codes" agent-tunnel = { path = "../crates/agent-tunnel", features = ["test-utils"] } agent-tunnel-libsql = { path = "../crates/agent-tunnel-libsql" } agent-tunnel-proto = { path = "../crates/agent-tunnel-proto", features = ["serde"] } +axum = { version = "0.8", default-features = false, features = ["http1", "json", "tokio"] } base64 = "0.23" camino = "1" +devolutions-gateway-ai.path = "../crates/devolutions-gateway-ai" devolutions-gateway-task = { path = "../crates/devolutions-gateway-task" } devolutions-gateway = { path = "../devolutions-gateway" } futures-util = "0.3" @@ -47,6 +49,7 @@ mcp-proxy.path = "../crates/mcp-proxy" network-scanner = { path = "../crates/network-scanner", features = ["test-utils"] } network-scanner-proto = { path = "../crates/network-scanner-proto" } nonempty = "0.12" +parking_lot = "0.12" picky = { version = "7.0.0-rc.25", default-features = false, features = ["jose"] } proxy-socks = { path = "../crates/proxy-socks" } quinn = "0.11" @@ -62,6 +65,7 @@ sysevent-codes.path = "../crates/sysevent-codes" tempfile = "3" test-utils.path = "../crates/test-utils" tokio-rustls = { version = "0.26", features = ["ring"] } +url = "2.5" uuid = { version = "1", features = ["v4"] } [target.'cfg(unix)'.dev-dependencies] diff --git a/testsuite/tests/gateway_ai.rs b/testsuite/tests/gateway_ai.rs new file mode 100644 index 000000000..6e04625fa --- /dev/null +++ b/testsuite/tests/gateway_ai.rs @@ -0,0 +1,493 @@ +use std::sync::Arc; +use std::time::Duration; + +use axum::Router; +use axum::http::{HeaderMap, StatusCode, Uri}; +use axum::routing::post; +use devolutions_gateway_ai::session_actions::Action; +use devolutions_gateway_ai::{AiClient, BuildError, Error, Provider, Usage}; +use parking_lot::Mutex; +use tokio::net::TcpListener; +use url::Url; + +const API_KEY: &str = "sk-test-secret"; +const MODEL: &str = "model-under-test"; +const REPORTED_MODEL: &str = "model-under-test-2026-09-30"; +const USAGE: Usage = Usage { + input_tokens: 10, + output_tokens: 20, +}; +const ANSWER: &str = "{\"offsetSeconds\":1.5,\"description\":\"Listed files\",\"object\":\"/var/log\",\"parameters\":{\"Command\":\"ls\"}}\nnot an action\n{\"offsetSeconds\":4,\"description\":\"Opened a shell\"}"; + +#[derive(Debug)] +struct CapturedRequest { + path: String, + headers: HeaderMap, + body: serde_json::Value, +} + +async fn spawn_provider(status: StatusCode, response: serde_json::Value) -> (Url, Arc>>) { + let captured = Arc::new(Mutex::new(None)); + + let app = Router::new().fallback(post({ + let captured = Arc::clone(&captured); + move |uri: Uri, headers: HeaderMap, body: String| { + let captured = Arc::clone(&captured); + async move { + *captured.lock() = Some(CapturedRequest { + path: uri.path().to_owned(), + headers, + body: serde_json::from_str(&body).unwrap(), + }); + (status, axum::Json(response)) + } + } + })); + + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + tokio::spawn(async move { axum::serve(listener, app).await.unwrap() }); + + (Url::parse(&format!("http://{addr}/v1/")).unwrap(), captured) +} + +fn http_client() -> reqwest::Client { + reqwest::Client::builder().no_proxy().build().unwrap() +} + +fn client(provider: Provider, base_url: Url) -> AiClient { + AiClient::builder() + .provider(provider) + .model(MODEL) + .api_key(API_KEY) + .base_url(base_url) + .http_client(http_client()) + .build() + .unwrap() +} + +fn openai_response(content: &str) -> serde_json::Value { + serde_json::json!({ + "id": "chatcmpl-1", + "object": "chat.completion", + "created": 0, + "model": REPORTED_MODEL, + "choices": [{ + "index": 0, + "message": { "role": "assistant", "content": content }, + "finish_reason": "stop" + }], + "usage": { "prompt_tokens": 10, "completion_tokens": 20, "total_tokens": 30 } + }) +} + +fn anthropic_response(text: &str) -> serde_json::Value { + serde_json::json!({ + "id": "msg_1", + "type": "message", + "role": "assistant", + "model": REPORTED_MODEL, + "content": [{ "type": "text", "text": text }], + "stop_reason": "end_turn", + "stop_sequence": null, + "usage": { "input_tokens": 10, "output_tokens": 20 } + }) +} + +fn assert_parsed_actions(actions: &[Action]) { + assert_eq!(actions.len(), 2); + assert_eq!(actions[0].offset, Duration::from_millis(1500)); + assert_eq!(actions[0].description, "Listed files"); + assert_eq!(actions[0].object.as_deref(), Some("/var/log")); + assert_eq!(actions[0].parameters.get("Command").map(String::as_str), Some("ls")); + assert_eq!(actions[1].offset, Duration::from_secs(4)); + assert_eq!(actions[1].object, None); +} + +fn header<'a>(headers: &'a HeaderMap, name: &str) -> Option<&'a str> { + headers.get(name).map(|value| value.to_str().unwrap()) +} + +fn body_contains(body: &serde_json::Value, needle: &str) -> bool { + body.to_string().contains(needle) +} + +#[tokio::test] +async fn openai_chat_request_and_response() { + let (base_url, captured) = spawn_provider(StatusCode::OK, openai_response(ANSWER)).await; + + let response = client(Provider::OpenAi, base_url) + .describe_session_actions("[1.5] ls") + .max_output_tokens(1234) + .send() + .await + .unwrap(); + + assert_parsed_actions(&response.output); + assert_eq!(response.model.as_deref(), Some(REPORTED_MODEL)); + assert_eq!(response.usage, Some(USAGE)); + + let request = captured.lock().take().unwrap(); + assert_eq!(request.path, "/v1/chat/completions"); + assert_eq!( + header(&request.headers, "authorization"), + Some(format!("Bearer {API_KEY}").as_str()) + ); + assert_eq!(header(&request.headers, "x-api-key"), None); + assert_eq!(request.body["model"], MODEL); + assert!(body_contains(&request.body, "[1.5] ls")); + assert_eq!(request.body["max_completion_tokens"], 1234); +} + +#[tokio::test] +async fn anthropic_messages_request_and_response() { + let (base_url, captured) = spawn_provider(StatusCode::OK, anthropic_response(ANSWER)).await; + + let response = client(Provider::Anthropic, base_url) + .describe_session_actions("[1.5] ls") + .max_output_tokens(1234) + .send() + .await + .unwrap(); + + assert_parsed_actions(&response.output); + assert_eq!(response.model.as_deref(), Some(REPORTED_MODEL)); + assert_eq!(response.usage, Some(USAGE)); + + let request = captured.lock().take().unwrap(); + assert_eq!(request.path, "/v1/messages"); + assert_eq!(header(&request.headers, "x-api-key"), Some(API_KEY)); + assert_eq!(header(&request.headers, "anthropic-version"), Some("2023-06-01")); + assert_eq!(header(&request.headers, "authorization"), None); + assert_eq!(request.body["model"], MODEL); + assert_eq!(request.body["max_tokens"], 1234); + assert!(body_contains(&request.body, "[1.5] ls")); +} + +#[tokio::test] +async fn provider_error_is_redacted_and_permanent() { + let body = serde_json::json!({ "error": { "message": format!("Incorrect API key provided: {API_KEY}") } }); + let (base_url, _captured) = spawn_provider(StatusCode::UNAUTHORIZED, body).await; + + let error = client(Provider::OpenAi, base_url) + .describe_session_actions("[0] whoami") + .send() + .await + .unwrap_err(); + + assert!(matches!(error, Error::Status { status: 401, .. }), "{error:?}"); + assert!(!error.is_transient()); + assert!(!error.to_string().contains(API_KEY), "{error}"); + assert!(!format!("{error:?}").contains(API_KEY), "{error:?}"); +} + +#[tokio::test] +async fn rate_limit_is_transient() { + let body = serde_json::json!({ "error": { "message": "Rate limit reached" } }); + let (base_url, _captured) = spawn_provider(StatusCode::TOO_MANY_REQUESTS, body).await; + + let error = client(Provider::Anthropic, base_url) + .describe_session_actions("[0] whoami") + .send() + .await + .unwrap_err(); + + assert!(matches!(error, Error::Status { status: 429, .. }), "{error:?}"); + assert!(error.is_transient()); +} + +#[tokio::test] +async fn unreachable_provider_is_transient() { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + drop(listener); + + let error = client(Provider::OpenAi, Url::parse(&format!("http://{addr}/v1/")).unwrap()) + .describe_session_actions("[0] whoami") + .send() + .await + .unwrap_err(); + + assert!(matches!(error, Error::Transport { .. }), "{error:?}"); + assert!(error.is_transient()); +} + +#[tokio::test] +async fn openai_compatible_uses_the_given_base_url() { + let (base_url, captured) = spawn_provider(StatusCode::OK, openai_response(ANSWER)).await; + + let response = client(Provider::OpenAiCompatible, base_url) + .describe_session_actions("[1.5] ls") + .send() + .await + .unwrap(); + + assert_parsed_actions(&response.output); + + let request = captured.lock().take().unwrap(); + assert_eq!(request.path, "/v1/chat/completions"); + assert_eq!( + header(&request.headers, "authorization"), + Some(format!("Bearer {API_KEY}").as_str()) + ); + assert_eq!(request.body["max_tokens"], 16_000); + assert!(request.body.get("max_completion_tokens").is_none()); +} + +#[tokio::test] +async fn gemini_speaks_openai_chat_with_max_tokens() { + let (base_url, captured) = spawn_provider(StatusCode::OK, openai_response(ANSWER)).await; + + let response = client(Provider::Gemini, base_url) + .describe_session_actions("[1.5] ls") + .max_output_tokens(1234) + .send() + .await + .unwrap(); + + assert_parsed_actions(&response.output); + + let request = captured.lock().take().unwrap(); + assert_eq!(request.path, "/v1/chat/completions"); + assert_eq!( + header(&request.headers, "authorization"), + Some(format!("Bearer {API_KEY}").as_str()) + ); + assert_eq!(request.body["max_tokens"], 1234); + assert!(request.body.get("max_completion_tokens").is_none()); +} + +#[tokio::test] +async fn think_blocks_are_not_read_as_actions() { + let answer = format!("\n{{\"offsetSeconds\":9,\"description\":\"Drafted an action\"}}\n\n{ANSWER}"); + let (base_url, _captured) = spawn_provider(StatusCode::OK, openai_response(&answer)).await; + + let response = client(Provider::OpenAiCompatible, base_url) + .describe_session_actions("[1.5] ls") + .send() + .await + .unwrap(); + + assert_parsed_actions(&response.output); +} + +#[tokio::test] +async fn silent_provider_times_out_as_transient() { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + // Accepts connections and never answers. + tokio::spawn(async move { + loop { + let (stream, _) = listener.accept().await.unwrap(); + tokio::spawn(async move { + let _open = stream; + std::future::pending::<()>().await; + }); + } + }); + + let started = std::time::Instant::now(); + let error = AiClient::builder() + .provider(Provider::OpenAiCompatible) + .model(MODEL) + .api_key(API_KEY) + .base_url(Url::parse(&format!("http://{addr}/v1/")).unwrap()) + .http_client(http_client()) + .request_timeout(Duration::from_millis(200)) + .build() + .unwrap() + .describe_session_actions("[0] whoami") + .send() + .await + .unwrap_err(); + + assert!(matches!(error, Error::Transport { .. }), "{error:?}"); + assert!(error.is_transient()); + assert!(started.elapsed() < Duration::from_secs(5), "{:?}", started.elapsed()); +} + +#[tokio::test] +async fn answer_without_model_or_usage_is_accepted() { + let mut answer = openai_response(ANSWER); + let fields = answer.as_object_mut().unwrap(); + fields.remove("model"); + fields.remove("usage"); + let (base_url, _captured) = spawn_provider(StatusCode::OK, answer).await; + + let response = client(Provider::OpenAiCompatible, base_url) + .describe_session_actions("[1.5] ls") + .send() + .await + .unwrap(); + + assert_parsed_actions(&response.output); + assert_eq!(response.model, None); + assert_eq!(response.usage, None); +} + +#[tokio::test] +async fn answers_cut_at_the_token_limit_are_truncated() { + let mut openai = openai_response(ANSWER); + openai["choices"][0]["finish_reason"] = "length".into(); + let mut anthropic = anthropic_response(ANSWER); + anthropic["stop_reason"] = "max_tokens".into(); + + for (provider, response) in [(Provider::OpenAi, openai), (Provider::Anthropic, anthropic)] { + let (base_url, _captured) = spawn_provider(StatusCode::OK, response).await; + + let error = client(provider, base_url) + .describe_session_actions("[1.5] ls") + .send() + .await + .unwrap_err(); + + assert!( + matches!(error, Error::Truncated { usage: Some(USAGE) }), + "{provider:?}: {error:?}" + ); + assert!(!error.is_transient()); + } +} + +#[tokio::test] +async fn answer_without_any_action_line_is_invalid_output() { + let (base_url, _captured) = spawn_provider(StatusCode::OK, openai_response("Sorry, I cannot help.")).await; + + let error = client(Provider::OpenAi, base_url) + .describe_session_actions("[1.5] ls") + .send() + .await + .unwrap_err(); + + assert!(matches!(error, Error::InvalidOutput { .. }), "{error:?}"); + assert!(!error.is_transient()); +} + +fn complete_builder(provider: Provider) -> devolutions_gateway_ai::AiClientBuilder { + AiClient::builder() + .provider(provider) + .model(MODEL) + .api_key(API_KEY) + .base_url(Url::parse("http://127.0.0.1:1/v1/").unwrap()) + .http_client(http_client()) +} + +#[test] +fn build_requires_provider() { + let result = AiClient::builder() + .model(MODEL) + .api_key(API_KEY) + .http_client(http_client()) + .build(); + + assert!(matches!(result, Err(BuildError::MissingProvider)), "{result:?}"); +} + +#[test] +fn build_requires_model() { + let result = AiClient::builder() + .provider(Provider::OpenAi) + .api_key(API_KEY) + .http_client(http_client()) + .build(); + assert!(matches!(result, Err(BuildError::MissingModel)), "{result:?}"); + + let result = complete_builder(Provider::OpenAi).model(" ").build(); + assert!(matches!(result, Err(BuildError::MissingModel)), "{result:?}"); +} + +#[test] +fn build_requires_key_for_hosted_providers() { + for provider in [ + Provider::OpenAi, + Provider::Anthropic, + Provider::Mistral, + Provider::Gemini, + Provider::OpenAiCompatible, + ] { + let result = AiClient::builder() + .provider(provider) + .model(MODEL) + .base_url(Url::parse("http://127.0.0.1:1/v1/").unwrap()) + .http_client(http_client()) + .build(); + assert!( + matches!(result, Err(BuildError::MissingApiKey(missing)) if missing == provider), + "{provider:?}: {result:?}" + ); + + let result = complete_builder(provider).api_key("").build(); + assert!( + matches!(result, Err(BuildError::MissingApiKey(missing)) if missing == provider), + "{provider:?}: {result:?}" + ); + } +} + +#[test] +fn build_rejects_a_key_that_is_not_a_header_value() { + let result = complete_builder(Provider::Anthropic).api_key("sk-test\r\n").build(); + + assert!(matches!(result, Err(BuildError::InvalidApiKey)), "{result:?}"); +} + +#[test] +fn build_requires_base_url_without_default() { + let result = AiClient::builder() + .provider(Provider::OpenAiCompatible) + .model(MODEL) + .api_key(API_KEY) + .http_client(http_client()) + .build(); + + assert!( + matches!(result, Err(BuildError::MissingBaseUrl(Provider::OpenAiCompatible))), + "{result:?}" + ); +} + +#[test] +fn build_rejects_a_base_url_that_is_not_http() { + let result = complete_builder(Provider::OpenAiCompatible) + .base_url(Url::parse("file:///etc/").unwrap()) + .build(); + + assert!(matches!(result, Err(BuildError::UnsupportedBaseUrl)), "{result:?}"); +} + +#[test] +fn build_uses_default_base_url() { + for provider in [ + Provider::OpenAi, + Provider::Anthropic, + Provider::Mistral, + Provider::Gemini, + ] { + let result = AiClient::builder() + .provider(provider) + .model(MODEL) + .api_key(API_KEY) + .http_client(http_client()) + .build(); + + assert!(result.is_ok(), "{provider:?}: {result:?}"); + } +} + +#[test] +fn build_requires_http_client() { + let result = AiClient::builder() + .provider(Provider::OpenAi) + .model(MODEL) + .api_key(API_KEY) + .build(); + + assert!(matches!(result, Err(BuildError::MissingHttpClient)), "{result:?}"); +} + +#[test] +fn client_debug_does_not_show_key() { + let ai = complete_builder(Provider::Anthropic).build().unwrap(); + + assert!(!format!("{ai:?}").contains(API_KEY)); +} diff --git a/testsuite/tests/main.rs b/testsuite/tests/main.rs index 4b5861aa8..d8c3f3850 100644 --- a/testsuite/tests/main.rs +++ b/testsuite/tests/main.rs @@ -4,6 +4,7 @@ mod agent_tunnel; mod cli; +mod gateway_ai; mod mcp_proxy; mod network_scanner; #[cfg(windows)]