From 5c01ddc0f69c4253e61c6e23432fe04ff48515cd Mon Sep 17 00:00:00 2001 From: Junyi Ou Date: Tue, 29 Sep 2026 16:47:55 -0400 Subject: [PATCH 1/3] feat(dgw): add an AI crate for session action descriptions `devolutions-gateway-ai` asks a model which actions a user performed in a session transcript and parses its JSON Lines answer into typed actions. It talks to the providers' HTTP APIs directly, with the workspace reqwest client passed by the caller (so Gateway's proxy and TLS policy apply): OpenAI chat completions (OpenAI, Mistral, any OpenAI-compatible endpoint) and Anthropic Messages. Only the fields of a single text completion are modeled, so there is no AI framework dependency and no new crate in the lockfile. The API key is redacted from every error. Co-Authored-By: Claude Opus 5.5 (1M context) --- Cargo.lock | 17 + crates/devolutions-gateway-ai/Cargo.toml | 19 + crates/devolutions-gateway-ai/src/lib.rs | 699 +++++++++++++++++++++++ testsuite/Cargo.toml | 4 + testsuite/tests/gateway_ai.rs | 303 ++++++++++ testsuite/tests/main.rs | 1 + 6 files changed, 1043 insertions(+) create mode 100644 crates/devolutions-gateway-ai/Cargo.toml create mode 100644 crates/devolutions-gateway-ai/src/lib.rs create mode 100644 testsuite/tests/gateway_ai.rs 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..9379b878c --- /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 = "Purpose-level AI requests for Devolutions Gateway" +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/lib.rs b/crates/devolutions-gateway-ai/src/lib.rs new file mode 100644 index 000000000..4fe393660 --- /dev/null +++ b/crates/devolutions-gateway-ai/src/lib.rs @@ -0,0 +1,699 @@ +//! Purpose-level AI helpers for Devolutions Gateway. +//! +//! Each provider is reached through its own HTTP API: OpenAI chat completions (also spoken by Mistral and many +//! others) or Anthropic Messages. Only the few fields a single text completion needs are modeled. + +use std::collections::BTreeMap; +use std::fmt; +use std::time::Duration; + +pub use reqwest; +pub use secrecy; +use secrecy::{ExposeSecret as _, SecretString}; +use serde::{Deserialize, Serialize}; +use tracing::{debug, warn}; +use url::Url; + +/// Version of the prompt used by [`AiClient::describe_session_actions`]. +/// +/// Bump it whenever the session actions prompt changes, so readers of the results know which prompt produced them. +pub const PROMPT_VERSION: &str = "session-actions-1"; + +const SESSION_ACTIONS_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."#; + +const DEFAULT_MAX_OUTPUT_TOKENS: u32 = 4_096; + +const ANTHROPIC_VERSION: &str = "2023-06-01"; + +/// 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, + /// Any endpoint speaking OpenAI chat completions, such as Gemini; the base URL is required. + OpenAiCompatible, +} + +impl Provider { + fn default_base_url(self) -> Option<&'static str> { + match self { + Self::OpenAi => Some("https://api.openai.com/v1/"), + Self::Anthropic => Some("https://api.anthropic.com/v1/"), + Self::Mistral => Some("https://api.mistral.ai/v1/"), + Self::OpenAiCompatible => None, + } + } +} + +/// 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, +} + +/// 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), + #[error("base URL is missing for AI provider {0:?}")] + MissingBaseUrl(Provider), + #[error("HTTP client is missing")] + MissingHttpClient, +} + +/// Error returned by [`DescribeSessionActions::send`]. +#[derive(Debug, thiserror::Error)] +#[non_exhaustive] +pub enum Error { + /// The provider request failed; `message` never contains the API key. + /// + /// `status` is the HTTP status of the provider answer, when there is one. + /// The underlying error is not kept as `source()`, because its chain could expose the key unredacted. + #[error("AI provider request failed: {message}")] + Request { status: Option, message: String }, + /// The provider answered with a body that is not the expected format. + #[error("AI provider answer is not valid: {reason}")] + InvalidResponse { reason: String }, + /// The answer has lines, but none of them is a valid action. + #[error("AI response has no valid action line ({invalid_lines} invalid lines)")] + NoValidAction { invalid_lines: usize }, +} + +/// 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, +} + +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 the provider default; 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 + } + + 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))?; + + 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, + }) + } +} + +fn resolve_base_url(provider: Provider, base_url: Option) -> Result { + match (base_url, provider.default_base_url()) { + (Some(base_url), _) => Ok(base_url), + (None, Some(default)) => Ok(Url::parse(default).expect("default base URLs are valid")), + (None, None) => Err(BuildError::MissingBaseUrl(provider)), + } +} + +/// Runs purpose-level AI requests against one provider. +#[derive(Clone)] +pub struct AiClient { + provider: Provider, + model: String, + base_url: Url, + api_key: SecretString, + http_client: reqwest::Client, +} + +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()) + .finish_non_exhaustive() + } +} + +impl AiClient { + pub fn builder() -> AiClientBuilder { + AiClientBuilder::default() + } + + /// Asks the model which actions the user performed in a session transcript. + /// + /// Each line of `input` must start with the elapsed time in seconds between square brackets, such as `[12.5] ls`. + pub fn describe_session_actions<'a>(&'a self, input: &'a str) -> DescribeSessionActions<'a> { + DescribeSessionActions { + client: self, + input, + max_output_tokens: DEFAULT_MAX_OUTPUT_TOKENS, + } + } + + async fn complete(&self, system: &str, input: &str, max_output_tokens: u32) -> Result { + debug!( + provider = ?self.provider, + model = %self.model, + base_url = %self.base_url, + input_len = input.len(), + max_output_tokens, + "Send AI completion request" + ); + + let model = self.model.as_str(); + let api_key = self.api_key.expose_secret(); + + let request = match self.provider { + Provider::OpenAi | Provider::Mistral | Provider::OpenAiCompatible => { + // OpenAI's newer models only accept `max_completion_tokens`; other servers only know `max_tokens`. + let (max_completion_tokens, max_tokens) = match self.provider { + Provider::OpenAi => (Some(max_output_tokens), None), + _ => (None, Some(max_output_tokens)), + }; + + self.http_client + .post(self.endpoint("chat/completions")) + .bearer_auth(api_key) + .json(&ChatRequest { + model, + messages: [ + ChatMessage { + role: "system", + content: system, + }, + ChatMessage { + role: "user", + content: input, + }, + ], + max_completion_tokens, + max_tokens, + }) + } + Provider::Anthropic => self + .http_client + .post(self.endpoint("messages")) + .header("x-api-key", api_key) + .header("anthropic-version", ANTHROPIC_VERSION) + .json(&MessagesRequest { + model, + system, + messages: [ChatMessage { + role: "user", + content: input, + }], + max_tokens: max_output_tokens, + }), + }; + + let response = request.send().await.map_err(|error| self.request_error(None, &error))?; + + let status = response.status(); + let body = response + .bytes() + .await + .map_err(|error| self.request_error(Some(status.as_u16()), &error))?; + + if !status.is_success() { + let message = serde_json::from_slice::(&body) + .map(|body| body.error.message) + .unwrap_or_else(|_| status.to_string()); + + return Err(Error::Request { + status: Some(status.as_u16()), + message: redact(message, api_key), + }); + } + + let text = match self.provider { + Provider::OpenAi | Provider::Mistral | Provider::OpenAiCompatible => { + let response: ChatResponse = parse_response(&body)?; + response + .choices + .into_iter() + .next() + .and_then(|choice| choice.message.content) + .unwrap_or_default() + } + Provider::Anthropic => { + let response: MessagesResponse = parse_response(&body)?; + response + .content + .into_iter() + .filter_map(|block| match block { + ContentBlock::Text { text } => Some(text), + ContentBlock::Other => None, + }) + .collect::>() + .join("\n") + } + }; + + debug!(output_len = text.len(), "Received AI completion"); + + Ok(text) + } + + fn endpoint(&self, path: &str) -> String { + format!("{}/{path}", self.base_url.as_str().trim_end_matches('/')) + } + + fn request_error(&self, status: Option, error: &dyn std::error::Error) -> Error { + request_error(status, error, self.api_key.expose_secret()) + } +} + +/// Request built by [`AiClient::describe_session_actions`]. +#[must_use = "the request is sent only by `send`"] +pub struct DescribeSessionActions<'a> { + client: &'a AiClient, + input: &'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("input_len", &self.input.len()) + .field("max_output_tokens", &self.max_output_tokens) + .finish() + } +} + +impl DescribeSessionActions<'_> { + /// Upper bound of tokens in the answer; the default is 4096. + pub fn max_output_tokens(mut self, max_output_tokens: u32) -> Self { + self.max_output_tokens = max_output_tokens; + self + } + + /// Invalid lines in the model answer are skipped with a warning. + /// It is an error only when the answer has lines but none of them is a valid action. + pub async fn send(self) -> Result, Error> { + let answer = self + .client + .complete(SESSION_ACTIONS_PROMPT, self.input, self.max_output_tokens) + .await?; + parse_actions(&answer) + } +} + +#[derive(Serialize)] +struct ChatMessage<'a> { + role: &'static str, + content: &'a str, +} + +/// OpenAI chat completions request. +#[derive(Serialize)] +struct ChatRequest<'a> { + model: &'a str, + messages: [ChatMessage<'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 { + choices: Vec, +} + +#[derive(Deserialize)] +struct ChatChoice { + message: ChatAnswer, +} + +#[derive(Deserialize)] +struct ChatAnswer { + content: Option, +} + +/// Anthropic Messages request. +#[derive(Serialize)] +struct MessagesRequest<'a> { + model: &'a str, + system: &'a str, + messages: [ChatMessage<'a>; 1], + max_tokens: u32, +} + +#[derive(Deserialize)] +struct MessagesResponse { + content: Vec, +} + +#[derive(Deserialize)] +#[serde(tag = "type", rename_all = "snake_case")] +enum ContentBlock { + Text { + text: String, + }, + #[serde(other)] + Other, +} + +/// Error body shared by OpenAI-style and Anthropic APIs. +#[derive(Deserialize)] +struct ProviderErrorBody { + error: ProviderError, +} + +#[derive(Deserialize)] +struct ProviderError { + message: String, +} + +// The reason never quotes the body, because the answer may contain session data. +fn parse_response(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 request_error(status: Option, 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::Request { + status, + 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]") + } +} + +#[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::NoValidAction { 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 { + #![allow(clippy::unwrap_used, reason = "test code can panic on errors")] + + 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).unwrap(); + + 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).unwrap(); + + assert_eq!(actions.len(), 1); + assert_eq!(actions[0].description, "Listed files"); + } + + #[test] + fn empty_answer_means_no_action() { + assert_eq!(parse_actions("").unwrap(), Vec::new()); + assert_eq!(parse_actions("\n```\n```\n").unwrap(), Vec::new()); + } + + #[test] + fn answer_without_valid_line_is_an_error() { + let error = parse_actions("not json\n{\"description\":\"no offset\"}\n").unwrap_err(); + + assert!(matches!(error, Error::NoValidAction { invalid_lines: 2 })); + } + + #[test] + fn invalid_line_reason_does_not_quote_the_line() { + let reason = + parse_action_line("{\"offsetSeconds\":1,\"description\":\"secret-value\",\"parameters\":{\"a\":1}}") + .unwrap_err(); + + assert!(!reason.contains("secret-value")); + } + + #[test] + fn invalid_response_reason_does_not_quote_the_body() { + let error = parse_response::(br#"{"choices":"secret-value"}"#) + .err() + .unwrap(); + + assert!(matches!(error, Error::InvalidResponse { .. })); + assert!(!error.to_string().contains("secret-value")); + } + + #[test] + fn prompt_asks_for_the_parsed_fields() { + for field in ["offsetSeconds", "description", "object", "parameters", "JSON Lines"] { + assert!(SESSION_ACTIONS_PROMPT.contains(field), "prompt is missing {field}"); + } + } + + #[test] + fn debug_redacts_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 describe_session_actions_debug_hides_input() { + let client = AiClient::builder() + .provider(Provider::OpenAi) + .model("gpt-test") + .api_key("sk-very-secret") + .http_client(reqwest::Client::new()) + .build() + .unwrap(); + + let debug = format!("{:?}", client.describe_session_actions("[1] secret-command")); + + assert!(!debug.contains("secret-command")); + assert!(debug.contains("input_len: 18")); + } + + #[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/"), + ] { + assert_eq!(resolve_base_url(provider, None).unwrap().as_str(), expected); + } + + assert!(matches!( + resolve_base_url(Provider::OpenAiCompatible, None), + Err(BuildError::MissingBaseUrl(Provider::OpenAiCompatible)) + )); + } + + #[test] + fn base_url_overrides_default() { + let custom = Url::parse("https://proxy.example/anthropic/").unwrap(); + + assert_eq!( + resolve_base_url(Provider::Anthropic, Some(custom.clone())).unwrap(), + custom + ); + } + + #[test] + fn error_message_redacts_api_key() { + let error = std::io::Error::other("invalid key sk-very-secret provided"); + + let error = request_error(Some(401), &error, "sk-very-secret"); + + assert!(!error.to_string().contains("sk-very-secret")); + assert!(!format!("{error:?}").contains("sk-very-secret")); + assert!(error.to_string().contains("[REDACTED]")); + } +} 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..1a3192858 --- /dev/null +++ b/testsuite/tests/gateway_ai.rs @@ -0,0 +1,303 @@ +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::{AiClient, BuildError, Provider}; +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 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": 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": 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: &[devolutions_gateway_ai::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 actions = client(Provider::OpenAi, base_url) + .describe_session_actions("[1.5] ls") + .max_output_tokens(1234) + .send() + .await + .unwrap(); + + assert_parsed_actions(&actions); + + 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 actions = client(Provider::Anthropic, base_url) + .describe_session_actions("[1.5] ls") + .max_output_tokens(1234) + .send() + .await + .unwrap(); + + assert_parsed_actions(&actions); + + 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() { + 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, devolutions_gateway_ai::Error::Request { status: Some(401), .. }), + "{error:?}" + ); + assert!(!error.to_string().contains(API_KEY), "{error}"); + assert!(!format!("{error:?}").contains(API_KEY), "{error:?}"); +} + +#[tokio::test] +async fn openai_compatible_uses_the_given_base_url() { + let (base_url, captured) = spawn_provider(StatusCode::OK, openai_response(ANSWER)).await; + + let actions = client(Provider::OpenAiCompatible, base_url) + .describe_session_actions("[1.5] ls") + .send() + .await + .unwrap(); + + assert_parsed_actions(&actions); + + 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"], 4096); + assert!(request.body.get("max_completion_tokens").is_none()); +} + +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::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_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_uses_default_base_url() { + for provider in [Provider::OpenAi, Provider::Anthropic, Provider::Mistral] { + 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)] From 017745d3f3cbd8f7d08215a6e21ae59eef9b92bb Mon Sep 17 00:00:00 2001 From: Junyi Ou Date: Wed, 30 Sep 2026 18:15:53 -0400 Subject: [PATCH 2/3] refactor(dgw): split the AI crate into purposes The crate had one purpose, and its shared parts were shaped by it: one prompt version for the whole crate, a purpose-specific error next to the transport ones, and retry rules left to every caller. Upcoming AI tasks need the same client with other prompts. Each purpose is now a module owning its prompt, prompt version and parser, and adding a method to AiClient; session_actions is the first. The provider formats live under wire/. Every purpose returns a Response with the output, the model reported by the provider and the token usage, which later budget and audit work needs. Error tells whether a request is worth retrying, reports answers cut at the output token limit as Truncated, and treats unusable model text as InvalidOutput. The builder rejects API keys that are not valid header values and base URLs that are not HTTP, so sending never fails on settings. --- crates/devolutions-gateway-ai/Cargo.toml | 2 +- crates/devolutions-gateway-ai/src/client.rs | 317 ++++++++ crates/devolutions-gateway-ai/src/error.rs | 146 ++++ crates/devolutions-gateway-ai/src/lib.rs | 710 +----------------- crates/devolutions-gateway-ai/src/response.rs | 46 ++ .../src/session_actions.rs | 272 +++++++ .../src/wire/anthropic.rs | 121 +++ crates/devolutions-gateway-ai/src/wire/mod.rs | 109 +++ .../devolutions-gateway-ai/src/wire/openai.rs | 143 ++++ testsuite/tests/gateway_ai.rs | 142 +++- 10 files changed, 1301 insertions(+), 707 deletions(-) create mode 100644 crates/devolutions-gateway-ai/src/client.rs create mode 100644 crates/devolutions-gateway-ai/src/error.rs create mode 100644 crates/devolutions-gateway-ai/src/response.rs create mode 100644 crates/devolutions-gateway-ai/src/session_actions.rs create mode 100644 crates/devolutions-gateway-ai/src/wire/anthropic.rs create mode 100644 crates/devolutions-gateway-ai/src/wire/mod.rs create mode 100644 crates/devolutions-gateway-ai/src/wire/openai.rs diff --git a/crates/devolutions-gateway-ai/Cargo.toml b/crates/devolutions-gateway-ai/Cargo.toml index 9379b878c..9ce0c17dc 100644 --- a/crates/devolutions-gateway-ai/Cargo.toml +++ b/crates/devolutions-gateway-ai/Cargo.toml @@ -3,7 +3,7 @@ name = "devolutions-gateway-ai" version = "0.0.0" edition = "2024" authors = ["Devolutions Inc. "] -description = "Purpose-level AI requests for Devolutions Gateway" +description = "AI requests for Devolutions Gateway, one method per purpose" publish = false [lints] diff --git a/crates/devolutions-gateway-ai/src/client.rs b/crates/devolutions-gateway-ai/src/client.rs new file mode 100644 index 000000000..8aefcb98b --- /dev/null +++ b/crates/devolutions-gateway-ai/src/client.rs @@ -0,0 +1,317 @@ +use std::fmt; + +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}; + +/// 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, + /// Any endpoint speaking OpenAI chat completions, such as Gemini; 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::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::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, +} + +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 + } + + /// 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, + }) + } +} + +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, +} + +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()) + .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. + 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 + .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)?, + }; + + debug!( + output_len = completion.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: completion.text, + model: completion.model, + usage: completion.usage, + }) + } +} + +#[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/"), + ] { + 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 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 index 4fe393660..9aa34f78e 100644 --- a/crates/devolutions-gateway-ai/src/lib.rs +++ b/crates/devolutions-gateway-ai/src/lib.rs @@ -1,699 +1,27 @@ -//! Purpose-level AI helpers for Devolutions Gateway. +//! 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 and many //! others) or Anthropic Messages. Only the few fields a single text completion needs are modeled. +//! +//! 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. -use std::collections::BTreeMap; -use std::fmt; -use std::time::Duration; +mod client; +mod error; +mod response; +pub mod session_actions; +mod wire; pub use reqwest; pub use secrecy; -use secrecy::{ExposeSecret as _, SecretString}; -use serde::{Deserialize, Serialize}; -use tracing::{debug, warn}; -use url::Url; - -/// Version of the prompt used by [`AiClient::describe_session_actions`]. -/// -/// Bump it whenever the session actions prompt changes, so readers of the results know which prompt produced them. -pub const PROMPT_VERSION: &str = "session-actions-1"; - -const SESSION_ACTIONS_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."#; - -const DEFAULT_MAX_OUTPUT_TOKENS: u32 = 4_096; - -const ANTHROPIC_VERSION: &str = "2023-06-01"; - -/// 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, - /// Any endpoint speaking OpenAI chat completions, such as Gemini; the base URL is required. - OpenAiCompatible, -} - -impl Provider { - fn default_base_url(self) -> Option<&'static str> { - match self { - Self::OpenAi => Some("https://api.openai.com/v1/"), - Self::Anthropic => Some("https://api.anthropic.com/v1/"), - Self::Mistral => Some("https://api.mistral.ai/v1/"), - Self::OpenAiCompatible => None, - } - } -} - -/// 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, -} - -/// 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), - #[error("base URL is missing for AI provider {0:?}")] - MissingBaseUrl(Provider), - #[error("HTTP client is missing")] - MissingHttpClient, -} - -/// Error returned by [`DescribeSessionActions::send`]. -#[derive(Debug, thiserror::Error)] -#[non_exhaustive] -pub enum Error { - /// The provider request failed; `message` never contains the API key. - /// - /// `status` is the HTTP status of the provider answer, when there is one. - /// The underlying error is not kept as `source()`, because its chain could expose the key unredacted. - #[error("AI provider request failed: {message}")] - Request { status: Option, message: String }, - /// The provider answered with a body that is not the expected format. - #[error("AI provider answer is not valid: {reason}")] - InvalidResponse { reason: String }, - /// The answer has lines, but none of them is a valid action. - #[error("AI response has no valid action line ({invalid_lines} invalid lines)")] - NoValidAction { invalid_lines: usize }, -} - -/// 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, -} - -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 the provider default; 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 - } - - 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))?; - - 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, - }) - } -} - -fn resolve_base_url(provider: Provider, base_url: Option) -> Result { - match (base_url, provider.default_base_url()) { - (Some(base_url), _) => Ok(base_url), - (None, Some(default)) => Ok(Url::parse(default).expect("default base URLs are valid")), - (None, None) => Err(BuildError::MissingBaseUrl(provider)), - } -} - -/// Runs purpose-level AI requests against one provider. -#[derive(Clone)] -pub struct AiClient { - provider: Provider, - model: String, - base_url: Url, - api_key: SecretString, - http_client: reqwest::Client, -} - -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()) - .finish_non_exhaustive() - } -} - -impl AiClient { - pub fn builder() -> AiClientBuilder { - AiClientBuilder::default() - } - - /// Asks the model which actions the user performed in a session transcript. - /// - /// Each line of `input` must start with the elapsed time in seconds between square brackets, such as `[12.5] ls`. - pub fn describe_session_actions<'a>(&'a self, input: &'a str) -> DescribeSessionActions<'a> { - DescribeSessionActions { - client: self, - input, - max_output_tokens: DEFAULT_MAX_OUTPUT_TOKENS, - } - } - - async fn complete(&self, system: &str, input: &str, max_output_tokens: u32) -> Result { - debug!( - provider = ?self.provider, - model = %self.model, - base_url = %self.base_url, - input_len = input.len(), - max_output_tokens, - "Send AI completion request" - ); - - let model = self.model.as_str(); - let api_key = self.api_key.expose_secret(); - - let request = match self.provider { - Provider::OpenAi | Provider::Mistral | Provider::OpenAiCompatible => { - // OpenAI's newer models only accept `max_completion_tokens`; other servers only know `max_tokens`. - let (max_completion_tokens, max_tokens) = match self.provider { - Provider::OpenAi => (Some(max_output_tokens), None), - _ => (None, Some(max_output_tokens)), - }; - - self.http_client - .post(self.endpoint("chat/completions")) - .bearer_auth(api_key) - .json(&ChatRequest { - model, - messages: [ - ChatMessage { - role: "system", - content: system, - }, - ChatMessage { - role: "user", - content: input, - }, - ], - max_completion_tokens, - max_tokens, - }) - } - Provider::Anthropic => self - .http_client - .post(self.endpoint("messages")) - .header("x-api-key", api_key) - .header("anthropic-version", ANTHROPIC_VERSION) - .json(&MessagesRequest { - model, - system, - messages: [ChatMessage { - role: "user", - content: input, - }], - max_tokens: max_output_tokens, - }), - }; - - let response = request.send().await.map_err(|error| self.request_error(None, &error))?; - - let status = response.status(); - let body = response - .bytes() - .await - .map_err(|error| self.request_error(Some(status.as_u16()), &error))?; - - if !status.is_success() { - let message = serde_json::from_slice::(&body) - .map(|body| body.error.message) - .unwrap_or_else(|_| status.to_string()); - - return Err(Error::Request { - status: Some(status.as_u16()), - message: redact(message, api_key), - }); - } - - let text = match self.provider { - Provider::OpenAi | Provider::Mistral | Provider::OpenAiCompatible => { - let response: ChatResponse = parse_response(&body)?; - response - .choices - .into_iter() - .next() - .and_then(|choice| choice.message.content) - .unwrap_or_default() - } - Provider::Anthropic => { - let response: MessagesResponse = parse_response(&body)?; - response - .content - .into_iter() - .filter_map(|block| match block { - ContentBlock::Text { text } => Some(text), - ContentBlock::Other => None, - }) - .collect::>() - .join("\n") - } - }; - - debug!(output_len = text.len(), "Received AI completion"); - - Ok(text) - } - - fn endpoint(&self, path: &str) -> String { - format!("{}/{path}", self.base_url.as_str().trim_end_matches('/')) - } - - fn request_error(&self, status: Option, error: &dyn std::error::Error) -> Error { - request_error(status, error, self.api_key.expose_secret()) - } -} - -/// Request built by [`AiClient::describe_session_actions`]. -#[must_use = "the request is sent only by `send`"] -pub struct DescribeSessionActions<'a> { - client: &'a AiClient, - input: &'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("input_len", &self.input.len()) - .field("max_output_tokens", &self.max_output_tokens) - .finish() - } -} - -impl DescribeSessionActions<'_> { - /// Upper bound of tokens in the answer; the default is 4096. - pub fn max_output_tokens(mut self, max_output_tokens: u32) -> Self { - self.max_output_tokens = max_output_tokens; - self - } - - /// Invalid lines in the model answer are skipped with a warning. - /// It is an error only when the answer has lines but none of them is a valid action. - pub async fn send(self) -> Result, Error> { - let answer = self - .client - .complete(SESSION_ACTIONS_PROMPT, self.input, self.max_output_tokens) - .await?; - parse_actions(&answer) - } -} - -#[derive(Serialize)] -struct ChatMessage<'a> { - role: &'static str, - content: &'a str, -} - -/// OpenAI chat completions request. -#[derive(Serialize)] -struct ChatRequest<'a> { - model: &'a str, - messages: [ChatMessage<'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 { - choices: Vec, -} - -#[derive(Deserialize)] -struct ChatChoice { - message: ChatAnswer, -} - -#[derive(Deserialize)] -struct ChatAnswer { - content: Option, -} - -/// Anthropic Messages request. -#[derive(Serialize)] -struct MessagesRequest<'a> { - model: &'a str, - system: &'a str, - messages: [ChatMessage<'a>; 1], - max_tokens: u32, -} - -#[derive(Deserialize)] -struct MessagesResponse { - content: Vec, -} - -#[derive(Deserialize)] -#[serde(tag = "type", rename_all = "snake_case")] -enum ContentBlock { - Text { - text: String, - }, - #[serde(other)] - Other, -} - -/// Error body shared by OpenAI-style and Anthropic APIs. -#[derive(Deserialize)] -struct ProviderErrorBody { - error: ProviderError, -} - -#[derive(Deserialize)] -struct ProviderError { - message: String, -} - -// The reason never quotes the body, because the answer may contain session data. -fn parse_response(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 request_error(status: Option, 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::Request { - status, - 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]") - } -} - -#[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::NoValidAction { 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 { - #![allow(clippy::unwrap_used, reason = "test code can panic on errors")] - - 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).unwrap(); - - 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).unwrap(); - - assert_eq!(actions.len(), 1); - assert_eq!(actions[0].description, "Listed files"); - } - - #[test] - fn empty_answer_means_no_action() { - assert_eq!(parse_actions("").unwrap(), Vec::new()); - assert_eq!(parse_actions("\n```\n```\n").unwrap(), Vec::new()); - } - - #[test] - fn answer_without_valid_line_is_an_error() { - let error = parse_actions("not json\n{\"description\":\"no offset\"}\n").unwrap_err(); - - assert!(matches!(error, Error::NoValidAction { invalid_lines: 2 })); - } - - #[test] - fn invalid_line_reason_does_not_quote_the_line() { - let reason = - parse_action_line("{\"offsetSeconds\":1,\"description\":\"secret-value\",\"parameters\":{\"a\":1}}") - .unwrap_err(); - - assert!(!reason.contains("secret-value")); - } - - #[test] - fn invalid_response_reason_does_not_quote_the_body() { - let error = parse_response::(br#"{"choices":"secret-value"}"#) - .err() - .unwrap(); - - assert!(matches!(error, Error::InvalidResponse { .. })); - assert!(!error.to_string().contains("secret-value")); - } - - #[test] - fn prompt_asks_for_the_parsed_fields() { - for field in ["offsetSeconds", "description", "object", "parameters", "JSON Lines"] { - assert!(SESSION_ACTIONS_PROMPT.contains(field), "prompt is missing {field}"); - } - } - - #[test] - fn debug_redacts_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 describe_session_actions_debug_hides_input() { - let client = AiClient::builder() - .provider(Provider::OpenAi) - .model("gpt-test") - .api_key("sk-very-secret") - .http_client(reqwest::Client::new()) - .build() - .unwrap(); - - let debug = format!("{:?}", client.describe_session_actions("[1] secret-command")); - - assert!(!debug.contains("secret-command")); - assert!(debug.contains("input_len: 18")); - } - - #[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/"), - ] { - assert_eq!(resolve_base_url(provider, None).unwrap().as_str(), expected); - } - - assert!(matches!( - resolve_base_url(Provider::OpenAiCompatible, None), - Err(BuildError::MissingBaseUrl(Provider::OpenAiCompatible)) - )); - } - - #[test] - fn base_url_overrides_default() { - let custom = Url::parse("https://proxy.example/anthropic/").unwrap(); - - assert_eq!( - resolve_base_url(Provider::Anthropic, Some(custom.clone())).unwrap(), - custom - ); - } - - #[test] - fn error_message_redacts_api_key() { - let error = std::io::Error::other("invalid key sk-very-secret provided"); - - let error = request_error(Some(401), &error, "sk-very-secret"); - assert!(!error.to_string().contains("sk-very-secret")); - assert!(!format!("{error:?}").contains("sk-very-secret")); - assert!(error.to_string().contains("[REDACTED]")); - } -} +pub use self::client::{AiClient, AiClientBuilder, BuildError, 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..da626ed39 --- /dev/null +++ b/crates/devolutions-gateway-ai/src/session_actions.rs @@ -0,0 +1,272 @@ +//! 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."#; + +const DEFAULT_MAX_OUTPUT_TOKENS: u32 = 4_096; + +/// 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 4096. + 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/tests/gateway_ai.rs b/testsuite/tests/gateway_ai.rs index 1a3192858..d26924d26 100644 --- a/testsuite/tests/gateway_ai.rs +++ b/testsuite/tests/gateway_ai.rs @@ -4,13 +4,19 @@ use std::time::Duration; use axum::Router; use axum::http::{HeaderMap, StatusCode, Uri}; use axum::routing::post; -use devolutions_gateway_ai::{AiClient, BuildError, Provider}; +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)] @@ -65,7 +71,7 @@ fn openai_response(content: &str) -> serde_json::Value { "id": "chatcmpl-1", "object": "chat.completion", "created": 0, - "model": MODEL, + "model": REPORTED_MODEL, "choices": [{ "index": 0, "message": { "role": "assistant", "content": content }, @@ -80,7 +86,7 @@ fn anthropic_response(text: &str) -> serde_json::Value { "id": "msg_1", "type": "message", "role": "assistant", - "model": MODEL, + "model": REPORTED_MODEL, "content": [{ "type": "text", "text": text }], "stop_reason": "end_turn", "stop_sequence": null, @@ -88,7 +94,7 @@ fn anthropic_response(text: &str) -> serde_json::Value { }) } -fn assert_parsed_actions(actions: &[devolutions_gateway_ai::Action]) { +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"); @@ -110,14 +116,16 @@ fn body_contains(body: &serde_json::Value, needle: &str) -> bool { async fn openai_chat_request_and_response() { let (base_url, captured) = spawn_provider(StatusCode::OK, openai_response(ANSWER)).await; - let actions = client(Provider::OpenAi, base_url) + let response = client(Provider::OpenAi, base_url) .describe_session_actions("[1.5] ls") .max_output_tokens(1234) .send() .await .unwrap(); - assert_parsed_actions(&actions); + 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"); @@ -135,14 +143,16 @@ async fn openai_chat_request_and_response() { async fn anthropic_messages_request_and_response() { let (base_url, captured) = spawn_provider(StatusCode::OK, anthropic_response(ANSWER)).await; - let actions = client(Provider::Anthropic, base_url) + let response = client(Provider::Anthropic, base_url) .describe_session_actions("[1.5] ls") .max_output_tokens(1234) .send() .await .unwrap(); - assert_parsed_actions(&actions); + 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"); @@ -155,7 +165,7 @@ async fn anthropic_messages_request_and_response() { } #[tokio::test] -async fn provider_error_is_redacted() { +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; @@ -165,25 +175,54 @@ async fn provider_error_is_redacted() { .await .unwrap_err(); - assert!( - matches!(error, devolutions_gateway_ai::Error::Request { status: Some(401), .. }), - "{error:?}" - ); + 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 actions = client(Provider::OpenAiCompatible, base_url) + let response = client(Provider::OpenAiCompatible, base_url) .describe_session_actions("[1.5] ls") .send() .await .unwrap(); - assert_parsed_actions(&actions); + assert_parsed_actions(&response.output); let request = captured.lock().take().unwrap(); assert_eq!(request.path, "/v1/chat/completions"); @@ -195,6 +234,63 @@ async fn openai_compatible_uses_the_given_base_url() { assert!(request.body.get("max_completion_tokens").is_none()); } +#[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) @@ -255,6 +351,13 @@ fn build_requires_key_for_hosted_providers() { } } +#[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() @@ -270,6 +373,15 @@ fn build_requires_base_url_without_default() { ); } +#[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] { From bb6b77f522dd82b9b5616d66eaa9ce6364090d24 Mon Sep 17 00:00:00 2001 From: Junyi Ou Date: Wed, 30 Sep 2026 19:48:50 -0400 Subject: [PATCH 3/3] feat(dgw): add Gemini and a request timeout to the AI crate Compared with how DVLS and RDM call AI providers, the crate had gaps. A request had no timeout, so a stalled provider held an AI task until the task timeout. The default output limit of 4096 tokens is low for reasoning models, which count their reasoning in it; DVLS and RDM use 16000 for Claude. Some models write a block before the answer, which DVLS removes. And Gemini, which DVLS offers, had no default base URL. Every request now has a timeout, 10 minutes by default and set on the builder, and a timeout is a transient error. The session actions default output limit is 16000. The blocks are removed from every answer before a purpose parses it. Provider::Gemini goes through Google's OpenAI-compatible endpoint. --- crates/devolutions-gateway-ai/src/client.rs | 86 ++++++++++++++++++- crates/devolutions-gateway-ai/src/lib.rs | 5 +- .../src/session_actions.rs | 5 +- testsuite/tests/gateway_ai.rs | 82 +++++++++++++++++- 4 files changed, 168 insertions(+), 10 deletions(-) diff --git a/crates/devolutions-gateway-ai/src/client.rs b/crates/devolutions-gateway-ai/src/client.rs index 8aefcb98b..d34bc3613 100644 --- a/crates/devolutions-gateway-ai/src/client.rs +++ b/crates/devolutions-gateway-ai/src/client.rs @@ -1,4 +1,5 @@ use std::fmt; +use std::time::Duration; use reqwest::header::HeaderValue; use secrecy::{ExposeSecret as _, SecretString}; @@ -8,6 +9,9 @@ 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 { @@ -17,7 +21,10 @@ pub enum Provider { Anthropic, /// Mistral chat completions; the default base URL is `https://api.mistral.ai/v1/`. Mistral, - /// Any endpoint speaking OpenAI chat completions, such as Gemini; the base URL is required. + /// 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, } @@ -28,6 +35,7 @@ impl Provider { 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, }; @@ -38,7 +46,7 @@ impl Provider { 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::OpenAiCompatible => Api::OpenAiChat(openai::TokenLimit::MaxTokens), + Self::Mistral | Self::Gemini | Self::OpenAiCompatible => Api::OpenAiChat(openai::TokenLimit::MaxTokens), Self::Anthropic => Api::AnthropicMessages, } } @@ -73,6 +81,7 @@ pub struct AiClientBuilder { api_key: Option, base_url: Option, http_client: Option, + request_timeout: Option, } impl AiClientBuilder { @@ -109,6 +118,15 @@ impl AiClientBuilder { 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)?; @@ -136,6 +154,7 @@ impl AiClientBuilder { base_url, api_key, http_client, + request_timeout: self.request_timeout.unwrap_or(DEFAULT_REQUEST_TIMEOUT), }) } } @@ -160,6 +179,7 @@ pub struct AiClient { base_url: Url, api_key: SecretString, http_client: reqwest::Client, + request_timeout: Duration, } impl fmt::Debug for AiClient { @@ -168,6 +188,7 @@ impl fmt::Debug for 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() } } @@ -189,6 +210,7 @@ impl AiClient { /// 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, @@ -212,6 +234,7 @@ impl AiClient { }; let response = request + .timeout(self.request_timeout) .send() .await .map_err(|error| error::transport(&error, api_key))?; @@ -231,8 +254,10 @@ impl AiClient { Api::AnthropicMessages => anthropic::parse(&body)?, }; + let text = strip_think_blocks(completion.text); + debug!( - output_len = completion.text.len(), + output_len = text.len(), truncated = completion.truncated, model = ?completion.model, usage = ?completion.usage, @@ -246,13 +271,48 @@ impl AiClient { } Ok(Response { - output: completion.text, + 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::*; @@ -273,6 +333,10 @@ mod tests { (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); @@ -284,6 +348,20 @@ mod tests { )); } + #[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"); diff --git a/crates/devolutions-gateway-ai/src/lib.rs b/crates/devolutions-gateway-ai/src/lib.rs index 9aa34f78e..3dddfae68 100644 --- a/crates/devolutions-gateway-ai/src/lib.rs +++ b/crates/devolutions-gateway-ai/src/lib.rs @@ -7,8 +7,9 @@ //! 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 and many +//! 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. @@ -22,6 +23,6 @@ mod wire; pub use reqwest; pub use secrecy; -pub use self::client::{AiClient, AiClientBuilder, BuildError, Provider}; +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/session_actions.rs b/crates/devolutions-gateway-ai/src/session_actions.rs index da626ed39..17f4d3b27 100644 --- a/crates/devolutions-gateway-ai/src/session_actions.rs +++ b/crates/devolutions-gateway-ai/src/session_actions.rs @@ -35,7 +35,8 @@ Rules: - Never copy passwords, secrets, or tokens. Write "[redacted]" instead. - If the user did nothing, write nothing."#; -const DEFAULT_MAX_OUTPUT_TOKENS: u32 = 4_096; +/// 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)] @@ -84,7 +85,7 @@ impl fmt::Debug for DescribeSessionActions<'_> { } impl DescribeSessionActions<'_> { - /// Upper bound of tokens in the answer; the default is 4096. + /// 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 diff --git a/testsuite/tests/gateway_ai.rs b/testsuite/tests/gateway_ai.rs index d26924d26..6e04625fa 100644 --- a/testsuite/tests/gateway_ai.rs +++ b/testsuite/tests/gateway_ai.rs @@ -230,10 +230,82 @@ async fn openai_compatible_uses_the_given_base_url() { header(&request.headers, "authorization"), Some(format!("Bearer {API_KEY}").as_str()) ); - assert_eq!(request.body["max_tokens"], 4096); + 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); @@ -330,6 +402,7 @@ fn build_requires_key_for_hosted_providers() { Provider::OpenAi, Provider::Anthropic, Provider::Mistral, + Provider::Gemini, Provider::OpenAiCompatible, ] { let result = AiClient::builder() @@ -384,7 +457,12 @@ fn build_rejects_a_base_url_that_is_not_http() { #[test] fn build_uses_default_base_url() { - for provider in [Provider::OpenAi, Provider::Anthropic, Provider::Mistral] { + for provider in [ + Provider::OpenAi, + Provider::Anthropic, + Provider::Mistral, + Provider::Gemini, + ] { let result = AiClient::builder() .provider(provider) .model(MODEL)