diff --git a/README.md b/README.md index 0c73fb4b7..1b39dc5c3 100644 --- a/README.md +++ b/README.md @@ -60,12 +60,12 @@ We hope in the future that we can find some way to upstream some / all of these ## Known Limitations -1. Semantic diff summarization is only supported for the Gemini class of models right now. Turn it on via: +1. Semantic diff summarization calls Gemini, OpenAI (or any OpenAI-compatible server, such as Ollama, via `endpoint`), or Anthropic. It is off by default. Turn it on via: a. Through the TUI: - Run `diffr config` - Search for 'summarization' - Enable in the dropdown - - Add your API key for Gemini if you don't already have on your path + - Pick a provider and model, and add its API key unless `GEMINI_API_KEY`, `OPENAI_API_KEY` or `ANTHROPIC_API_KEY` is already in your environment b. Through the [Whiteboard app](https://github.com/devdotfast/whiteboard). 2. The plugin API is a bit awkward and will be simplified radically in the coming releases. diff --git a/plugins/summarize/plugin.toml b/plugins/summarize/plugin.toml index b493f37f0..d92270543 100644 --- a/plugins/summarize/plugin.toml +++ b/plugins/summarize/plugin.toml @@ -7,13 +7,13 @@ description = "Pseudocode summaries for large new function bodies and tests." # startup when a plugin that is on cannot be made. [enabled] title = "Summarize functions and tests" -description = "Summarize large new function bodies and right-side tests as pseudocode. Needs an API key: `api_key`, or `GEMINI_API_KEY` or `GOOGLE_API_KEY` in the environment." +description = "Summarize large new function bodies and right-side tests as pseudocode. Needs an API key: `api_key`, or the provider's variable in the environment: `GEMINI_API_KEY` or `GOOGLE_API_KEY`, `OPENAI_API_KEY`, or `ANTHROPIC_API_KEY`." default = false [options.provider] -enum = ["gemini"] +enum = ["gemini", "openai", "anthropic"] title = "Provider" -description = "Which model API to call." +description = "Which model API to call. `openai` also reaches OpenRouter and any other OpenAI-compatible server through `endpoint`. Set `model` to one of the provider's models when you change this." default = "gemini" [options.model] @@ -32,12 +32,12 @@ default = 20 [options.api_key] type = "string" title = "API key" -description = "The provider's API key. `GEMINI_API_KEY` or `GOOGLE_API_KEY` in the environment is used when this is unset." +description = "Sent to the selected provider. The provider's variable in the environment is used when this is unset. Optional for `openai` with a custom `endpoint`." [options.endpoint] type = "string" title = "Endpoint URL" -description = "Override the provider's base URL, for proxies and tests." +description = "Override the provider's base URL, for proxies, OpenAI-compatible servers (for example `https://openrouter.ai/api/v1` or `http://localhost:11434/v1`) and tests. Empty uses the provider's default." [options.request_timeout_ms] type = "integer" @@ -71,8 +71,8 @@ For each listed fold, rewrite that function body as short \ pseudocode. Keep the names. No prose, no comments, no code fences. Use as few lines as \ possible: about one pseudocode line per five source lines, and never more than a third \ of the body's lines. When a fold lists a doc, also set "summary" to one sentence copied \ -verbatim from that doc; otherwise leave it empty. Answer with a JSON array of \ -{"id", "summary", "pseudocode"} objects, one per fold.""" +verbatim from that doc; otherwise leave it empty. Answer with a JSON object whose \ +"summaries" array holds one {"id", "summary", "pseudocode"} object per fold.""" [options.tests] type = "boolean" diff --git a/plugins/summarize/plugin.wasm b/plugins/summarize/plugin.wasm index 1a575b1c3..f5c7b3f5e 100644 Binary files a/plugins/summarize/plugin.wasm and b/plugins/summarize/plugin.wasm differ diff --git a/plugins/summarize/src/http.rs b/plugins/summarize/src/http.rs index cb30eba8b..4583cd037 100644 --- a/plugins/summarize/src/http.rs +++ b/plugins/summarize/src/http.rs @@ -4,7 +4,12 @@ use diffr_plugin_sdk::anyhow::{self, anyhow}; use wasi::http::{outgoing_handler, types::*}; use wasi::io::streams::StreamError; -pub fn post(url: &str, key: &str, body: &str, timeout_ms: u64) -> anyhow::Result<(u16, Vec)> { +pub fn post( + url: &str, + headers: &[(&'static str, String)], + body: &str, + timeout_ms: u64, +) -> anyhow::Result<(u16, Vec)> { let (scheme, rest) = url .split_once("://") .ok_or_else(|| anyhow!("invalid endpoint URL"))?; @@ -14,12 +19,19 @@ pub fn post(url: &str, key: &str, body: &str, timeout_ms: u64) -> anyhow::Result _ => anyhow::bail!("endpoint must use http or https"), }; let (authority, path) = rest.split_once('/').unwrap_or((rest, "")); - let headers = Fields::from_list(&[ - ("content-type".into(), b"application/json".to_vec()), - ("content-length".into(), body.len().to_string().into_bytes()), - ("x-goog-api-key".into(), key.as_bytes().to_vec()), - ]) - .map_err(|e| anyhow!("HTTP headers: {e:?}"))?; + let mut fields = vec![ + ("content-type".to_owned(), b"application/json".to_vec()), + ( + "content-length".to_owned(), + body.len().to_string().into_bytes(), + ), + ]; + fields.extend( + headers + .iter() + .map(|(name, value)| ((*name).to_owned(), value.as_bytes().to_vec())), + ); + let headers = Fields::from_list(&fields).map_err(|e| anyhow!("HTTP headers: {e:?}"))?; let request = OutgoingRequest::new(headers); request .set_method(&Method::Post) diff --git a/plugins/summarize/src/lib.rs b/plugins/summarize/src/lib.rs index f4abd5118..f9746717e 100644 --- a/plugins/summarize/src/lib.rs +++ b/plugins/summarize/src/lib.rs @@ -1,8 +1,9 @@ //! The summarizer: large new function bodies become short pseudocode, shown //! in place of the collapsed body. //! -//! It needs an API key: `new` fails without one, naming how to set it or -//! turn the plugin off, which is why the bundled configuration ships it off. +//! It needs an API key for most providers: `new` fails without one, naming +//! how to set it or turn the plugin off, which is why the bundled +//! configuration ships it off. //! Requests use WASI HTTP. The host calls one file at a time per instance. use diffr_plugin_sdk::anyhow::{self, anyhow, Context as _}; use diffr_plugin_sdk::{ @@ -10,10 +11,11 @@ use diffr_plugin_sdk::{ FileEntry, Move, Node, OtherSide, Pairing, Plugin, Region, Source, }; use serde::Deserialize; -use serde_json::json; use std::collections::BTreeMap; use std::time::Duration; mod http; +mod provider; +pub use provider::Provider; /// The plugin's name, and the tags its queries set: a function body, and a /// test body, which can be summarized independently of whether it is new. @@ -40,18 +42,10 @@ pub struct Options { pub system_prompt: String, } -#[derive(Clone, Copy, Debug, PartialEq, Eq, Deserialize)] -#[serde(rename_all = "lowercase")] -pub enum Provider { - Gemini, -} - -const DEFAULT_ENDPOINT: &str = "https://generativelanguage.googleapis.com"; - /// The summarizer's options, API key and endpoint. pub struct Summarize { options: Options, - api_key: String, + api_key: Option, endpoint: String, } @@ -80,16 +74,6 @@ struct Summary { } impl Summarize { - fn url(&self) -> String { - match self.options.provider { - Provider::Gemini => format!( - "{}/v1beta/models/{}:generateContent", - self.endpoint.trim_end_matches('/'), - self.options.model - ), - } - } - fn prompt(&self, path: &str, src: &str, folds: &[Request]) -> String { let numbered: Vec = src .split_terminator('\n') @@ -125,36 +109,22 @@ impl Summarize { src: &str, folds: &[Request], ) -> anyhow::Result> { - let body = json!({ - "systemInstruction": {"parts": [{"text": self.options.system_prompt}]}, - "contents": [{"role": "user", "parts": [{"text": self.prompt(path, src, folds)}]}], - "generationConfig": { - "temperature": 0, - "maxOutputTokens": 600 * folds.len() + 200, - "thinkingConfig": {"thinkingBudget": 0}, - "responseMimeType": "application/json", - "responseSchema": { - "type": "ARRAY", - "items": { - "type": "OBJECT", - "properties": { - "id": {"type": "INTEGER"}, - "summary": {"type": "STRING"}, - "pseudocode": {"type": "STRING"}, - }, - "required": ["id", "pseudocode"], - }, - }, - }, - }); - let url = self.url(); + let provider = self.options.provider; + let body = provider.body( + &self.options.model, + &self.options.system_prompt, + &self.prompt(path, src, folds), + 600 * folds.len() + 200, + ); + let url = provider.url(&self.endpoint, &self.options.model); + let headers = provider.headers(self.api_key.as_deref()); let failed = |message: String| anyhow!("{}: {message}", self.options.model); let text: serde_json::Value = { let mut attempt = 0; loop { let result = http::post( &url, - &self.api_key, + &headers, &body.to_string(), self.options.request_timeout_ms, ); @@ -182,13 +152,11 @@ impl Summarize { std::thread::sleep(Duration::from_millis(250 * (1 << attempt.min(6)))); } }; - let content = text["candidates"][0]["content"]["parts"] - .as_array() - .and_then(|parts| parts.last()) - .and_then(|part| part["text"].as_str()) + let content = provider + .text(&text) .ok_or_else(|| failed("no text in the response".to_owned()))?; - let answers: Vec = - serde_json::from_str(content).map_err(|error| failed(format!("{error}: {content}")))?; + let answers: Vec = provider::answers(content) + .ok_or_else(|| failed(format!("no summaries in the answer: {content}")))?; let mut texts = BTreeMap::new(); for answer in answers { if !folds.iter().any(|fold| fold.id == answer.id) { @@ -329,25 +297,31 @@ fn compresses(summary: &str, body: &[&str]) -> bool { summary_lines * 2 <= body_lines } -/// The API key: the `api_key` option, or else `GEMINI_API_KEY`, or else -/// `GOOGLE_API_KEY`, the first that is set and not empty. -fn resolve_key(config: &Options) -> anyhow::Result { +/// The API key: the `api_key` option, or else the first of the provider's +/// environment variables that is set and not empty. `None` only where the +/// provider can go without one. +fn resolve_key(config: &Options, custom_endpoint: bool) -> anyhow::Result> { let set = |key: &String| !key.is_empty(); if let Some(key) = config.api_key.clone().filter(set) { - return Ok(key); + return Ok(Some(key)); } - for variable in ["GEMINI_API_KEY", "GOOGLE_API_KEY"] { + let variables = config.provider.key_variables(); + for variable in variables { if let Some(key) = std::env::var_os(variable) { let key = key .into_string() .map_err(|_| anyhow!("{variable} is not valid UTF-8"))?; if set(&key) { - return Ok(key); + return Ok(Some(key)); } } } + if config.provider.key_optional(custom_endpoint) { + return Ok(None); + } anyhow::bail!( - "no API key: set plugins.bundled.summarize.api_key, or GEMINI_API_KEY or GOOGLE_API_KEY in the environment, or turn the summarizer off with plugins.bundled.summarize.enabled = false" + "no API key: set plugins.bundled.summarize.api_key, or {} in the environment, or turn the summarizer off with plugins.bundled.summarize.enabled = false", + variables.join(" or ") ) } @@ -355,13 +329,14 @@ impl Plugin for Summarize { type Options = Options; fn new(options: Options) -> anyhow::Result { - let api_key = resolve_key(&options)?; + let endpoint = options + .endpoint + .clone() + .filter(|endpoint| !endpoint.is_empty()); + let api_key = resolve_key(&options, endpoint.is_some())?; Ok(Self { api_key, - endpoint: options - .endpoint - .clone() - .unwrap_or_else(|| DEFAULT_ENDPOINT.to_owned()), + endpoint: endpoint.unwrap_or_else(|| options.provider.default_endpoint().to_owned()), options, }) } diff --git a/plugins/summarize/src/provider.rs b/plugins/summarize/src/provider.rs new file mode 100644 index 000000000..c741091b1 --- /dev/null +++ b/plugins/summarize/src/provider.rs @@ -0,0 +1,158 @@ +//! Each model API's wire format: where a request goes, how it is +//! authenticated, and where the answer's text is. +use serde::de::DeserializeOwned; +use serde::Deserialize; +use serde_json::{json, Value}; + +#[derive(Clone, Copy, Debug, PartialEq, Eq, Deserialize)] +#[serde(rename_all = "lowercase")] +pub enum Provider { + Gemini, + OpenAi, + Anthropic, +} + +impl Provider { + pub fn default_endpoint(self) -> &'static str { + match self { + Self::Gemini => "https://generativelanguage.googleapis.com", + Self::OpenAi => "https://api.openai.com/v1", + Self::Anthropic => "https://api.anthropic.com", + } + } + + /// The environment variables read, in order, when `api_key` is unset. + pub fn key_variables(self) -> &'static [&'static str] { + match self { + Self::Gemini => &["GEMINI_API_KEY", "GOOGLE_API_KEY"], + Self::OpenAi => &["OPENAI_API_KEY"], + Self::Anthropic => &["ANTHROPIC_API_KEY"], + } + } + + /// OpenAI-compatible servers at a custom endpoint, such as Ollama, + /// often take no key. + pub fn key_optional(self, custom_endpoint: bool) -> bool { + self == Self::OpenAi && custom_endpoint + } + + pub fn url(self, endpoint: &str, model: &str) -> String { + let endpoint = endpoint.trim_end_matches('/'); + match self { + Self::Gemini => format!("{endpoint}/v1beta/models/{model}:generateContent"), + Self::OpenAi => format!("{endpoint}/chat/completions"), + Self::Anthropic => format!("{endpoint}/v1/messages"), + } + } + + pub fn headers(self, key: Option<&str>) -> Vec<(&'static str, String)> { + let mut headers = Vec::new(); + if let Some(key) = key { + headers.push(match self { + Self::Gemini => ("x-goog-api-key", key.to_owned()), + Self::OpenAi => ("authorization", format!("Bearer {key}")), + Self::Anthropic => ("x-api-key", key.to_owned()), + }); + } + if self == Self::Anthropic { + headers.push(("anthropic-version", "2023-06-01".to_owned())); + } + headers + } + + /// Every provider constrains the answer with the same schema. + /// Only Gemini gets a temperature: current reasoning models reject one. + /// OpenAI gets no length limit either, since compatible servers name it + /// differently; Anthropic's limit also covers thinking, so it has a floor. + pub fn body(self, model: &str, system: &str, user: &str, max_tokens: usize) -> Value { + match self { + Self::Gemini => json!({ + "systemInstruction": {"parts": [{"text": system}]}, + "contents": [{"role": "user", "parts": [{"text": user}]}], + "generationConfig": { + "temperature": 0, + "maxOutputTokens": max_tokens, + "thinkingConfig": {"thinkingBudget": 0}, + "responseMimeType": "application/json", + "responseJsonSchema": summaries_schema(), + }, + }), + Self::OpenAi => json!({ + "model": model, + "messages": [ + {"role": "system", "content": system}, + {"role": "user", "content": user}, + ], + "response_format": { + "type": "json_schema", + "json_schema": {"name": "summaries", "strict": true, "schema": summaries_schema()}, + }, + }), + Self::Anthropic => json!({ + "model": model, + "max_tokens": max_tokens.max(4096), + "system": system, + "messages": [{"role": "user", "content": user}], + "output_config": {"format": {"type": "json_schema", "schema": summaries_schema()}}, + }), + } + } + + pub fn text(self, response: &Value) -> Option<&str> { + match self { + Self::Gemini => response["candidates"][0]["content"]["parts"] + .as_array()? + .last()?["text"] + .as_str(), + Self::OpenAi => response["choices"][0]["message"]["content"].as_str(), + Self::Anthropic => response["content"] + .as_array()? + .iter() + .rev() + .find(|block| block["type"] == "text")?["text"] + .as_str(), + } + } +} + +/// `{"summaries": [{id, summary, pseudocode}]}`, every field required: OpenAI +/// and Anthropic take only an object at the root. +fn summaries_schema() -> Value { + json!({ + "type": "object", + "properties": { + "summaries": { + "type": "array", + "items": { + "type": "object", + "properties": { + "id": {"type": "integer"}, + "summary": {"type": "string"}, + "pseudocode": {"type": "string"}, + }, + "required": ["id", "summary", "pseudocode"], + "additionalProperties": false, + }, + }, + }, + "required": ["summaries"], + "additionalProperties": false, + }) +} + +/// The first JSON array in an answer that holds items, ignoring any prose, +/// fence or reasoning around it, and the `summaries` object wrapping it; +/// an empty array only when none does. +pub fn answers(text: &str) -> Option> { + let mut parsed = text.match_indices('[').filter_map(|(start, _)| { + serde_json::Deserializer::from_str(&text[start..]) + .into_iter::>() + .next()? + .ok() + }); + let first = parsed.next()?; + if !first.is_empty() { + return Some(first); + } + Some(parsed.find(|items| !items.is_empty()).unwrap_or(first)) +} diff --git a/src/config/default.toml b/src/config/default.toml index 6056aa8a7..1dba1f8e7 100644 --- a/src/config/default.toml +++ b/src/config/default.toml @@ -46,7 +46,6 @@ test_min_lines = 20 request_timeout_ms = 60000 max_concurrency = 16 retries = 3 -system_prompt = """For each listed fold, rewrite that function body as short pseudocode. Keep the names. No prose, no comments, no code fences. Use as few lines as possible: about one pseudocode line per five source lines, and never more than a third of the body's lines. When a fold lists a doc, also set "summary" to one sentence copied verbatim from that doc; otherwise leave it empty. Answer with a JSON array of {"id", "summary", "pseudocode"} objects, one per fold.""" [plugins.bundled.test-bodies] enabled = true diff --git a/src/config/store.rs b/src/config/store.rs index 10ec5cc95..62139c316 100644 --- a/src/config/store.rs +++ b/src/config/store.rs @@ -309,8 +309,8 @@ mod tests { "plugins.bundled.context.lines: expected an integer, got \"many\"" ); assert_eq!( - error("plugins.bundled.summarize.provider", "openai"), - "plugins.bundled.summarize.provider: expected one of \"gemini\", got \"openai\"" + error("plugins.bundled.summarize.provider", "mistral"), + "plugins.bundled.summarize.provider: expected one of \"gemini\", \"openai\", \"anthropic\", got \"mistral\"" ); assert_eq!( error("plugins.external.mine.enabled", "true"), diff --git a/src/plugin/config.rs b/src/plugin/config.rs index 8293e3e05..17a0720f2 100644 --- a/src/plugin/config.rs +++ b/src/plugin/config.rs @@ -528,6 +528,20 @@ impl PluginsConfig { #[cfg(test)] mod tests { use crate::config::Config; + + #[test] + fn the_embedded_defaults_agree_with_each_plugin_toml() { + for (name, entry) in super::default_tables().bundled { + let defaults = super::builtin::manifest(&name) + .expect("a bundled plugin") + .defaults(); + for (key, value) in &entry.options { + if let Some(default) = defaults.get(key) { + assert_eq!(value, default, "{name}.{key}"); + } + } + } + } use serde_json::Value; /// Every setting the schema lists, as a settings screen flattens it: @@ -758,7 +772,7 @@ mod tests { zero.starts_with("plugins.bundled.summarize: max_concurrency: "), "{zero}" ); - assert!(error("[plugins.bundled.summarize]\nprovider = 'openai'\n") + assert!(error("[plugins.bundled.summarize]\nprovider = 'mistral'\n") .starts_with("plugins.bundled.summarize: provider: ")); } } diff --git a/src/plugin/tests/mod.rs b/src/plugin/tests/mod.rs index 2e6030813..1bc6f5500 100644 --- a/src/plugin/tests/mod.rs +++ b/src/plugin/tests/mod.rs @@ -442,6 +442,20 @@ fn a_plugin_that_cannot_be_made_is_a_setup_error() { format!("{error:#}"), "plugins.bundled.summarize: no API key: set plugins.bundled.summarize.api_key, or GEMINI_API_KEY or GOOGLE_API_KEY in the environment, or turn the summarizer off with plugins.bundled.summarize.enabled = false" ); + if std::env::var_os("OPENAI_API_KEY").is_some() { + return; + } + let config = Config::from_toml( + "[plugins.bundled.summarize]\nenabled = true\nprovider = 'openai'\nmodel = 'm'\napi_key = ''\nendpoint = ''\n", + ) + .unwrap(); + let error = Pipeline::from_config(&config.plugins, Path::new(".")) + .err() + .expect("OpenAI at its default endpoint needs a key"); + assert_eq!( + format!("{error:#}"), + "plugins.bundled.summarize: no API key: set plugins.bundled.summarize.api_key, or OPENAI_API_KEY in the environment, or turn the summarizer off with plugins.bundled.summarize.enabled = false" + ); } #[test] diff --git a/src/plugin/tests/summarize.rs b/src/plugin/tests/summarize.rs index 7b5e8f96a..25e9fa2b3 100644 --- a/src/plugin/tests/summarize.rs +++ b/src/plugin/tests/summarize.rs @@ -27,29 +27,47 @@ fn project_with( const LARGE: &str = "def f():\n a()\n b()\n c()\n\ndef g(): d()\n"; +/// One request the test server received. +struct Received { + line: String, + headers: Vec, + body: String, +} + /// Answer each request with the next canned response. -fn serve(responses: Vec<(u16, String)>) -> (String, std::thread::JoinHandle>) { +fn serve_requests( + responses: Vec<(u16, String)>, +) -> (String, std::thread::JoinHandle>) { let listener = TcpListener::bind("127.0.0.1:0").unwrap(); let endpoint = format!("http://{}", listener.local_addr().unwrap()); let handle = std::thread::spawn(move || { - let mut bodies = Vec::new(); + let mut received = Vec::new(); for (status, body) in responses { let (stream, _) = listener.accept().unwrap(); let mut reader = BufReader::new(stream); + let mut line = String::new(); + reader.read_line(&mut line).unwrap(); + let mut headers = Vec::new(); let mut length = 0; loop { - let mut line = String::new(); - reader.read_line(&mut line).unwrap(); - if line == "\r\n" { + let mut header = String::new(); + reader.read_line(&mut header).unwrap(); + if header == "\r\n" { break; } - if let Some(value) = line.to_ascii_lowercase().strip_prefix("content-length:") { + let header = header.trim_end().to_ascii_lowercase(); + if let Some(value) = header.strip_prefix("content-length:") { length = value.trim().parse().unwrap(); } + headers.push(header); } let mut request = vec![0; length]; reader.read_exact(&mut request).unwrap(); - bodies.push(String::from_utf8(request).unwrap()); + received.push(Received { + line: line.trim_end().to_owned(), + headers, + body: String::from_utf8(request).unwrap(), + }); let reason = if status == 200 { "OK" } else { "Error" }; let response = format!( "HTTP/1.1 {status} {reason}\r\nContent-Type: application/json\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{body}", @@ -57,11 +75,25 @@ fn serve(responses: Vec<(u16, String)>) -> (String, std::thread::JoinHandle) -> (String, std::thread::JoinHandle>) { + let (endpoint, handle) = serve_requests(responses); + let bodies = std::thread::spawn(move || { + handle + .join() + .unwrap() + .into_iter() + .map(|request| request.body) + .collect() + }); + (endpoint, bodies) +} + /// The label of the only function body on the after side. The scope fold /// the context queries wrap around it is not one. fn fold_label(sides: &tree::Pairing) -> String { @@ -80,8 +112,8 @@ fn gemini_answer(items: &[(u32, &str)]) -> String { .iter() .map(|(id, text)| json!({"id": id, "pseudocode": text})) .collect(); - json!({"candidates": [{"content": {"parts": [{"text": serde_json::to_string(&answers).unwrap()}]}}]}) - .to_string() + let text = json!({"summaries": answers}).to_string(); + json!({"candidates": [{"content": {"parts": [{"text": text}]}}]}).to_string() } fn summarizer(endpoint: &str, retries: u32) -> Pipeline { @@ -581,3 +613,147 @@ fn deferred_summary_preserves_user_fold_state_and_region_identity() { assert_eq!(sides, before); server.join().unwrap(); } + +#[test] +fn gemini_requests_keep_their_path_and_key_header() { + let (file, mut sides) = project("a.py", "", LARGE); + let id = select(&trees(&sides), 3, None)[0].0; + let (endpoint, server) = serve_requests(vec![(200, gemini_answer(&[(id, "call a, b, c")]))]); + summarizer(&endpoint, 0).run(&file, &mut sides).unwrap(); + let request = server.join().unwrap().remove(0); + assert_eq!( + request.line, + "POST /v1beta/models/gemini-3.8-flash:generateContent HTTP/1.1" + ); + assert!(request + .headers + .contains(&"x-goog-api-key: test-key".to_owned())); + let body: serde_json::Value = serde_json::from_str(&request.body).unwrap(); + let config = &body["generationConfig"]; + assert!(config.get("responseSchema").is_none()); + assert_eq!(config["responseMimeType"], "application/json"); + assert_summaries_schema(&config["responseJsonSchema"]); + assert!(!request + .headers + .iter() + .any(|header| header.starts_with("authorization"))); +} + +fn answers(items: &[(u32, &str)]) -> String { + let answers: Vec<_> = items + .iter() + .map(|(id, text)| json!({"id": id, "pseudocode": text})) + .collect(); + serde_json::to_string(&answers).unwrap() +} + +/// An object root with every field required and nothing else allowed, as +/// OpenAI's strict mode and Anthropic's structured outputs require. +fn assert_summaries_schema(schema: &serde_json::Value) { + assert_eq!(schema["type"], "object"); + assert_eq!(schema["required"], json!(["summaries"])); + assert_eq!(schema["additionalProperties"], false); + let item = &schema["properties"]["summaries"]["items"]; + assert_eq!(item["required"], json!(["id", "summary", "pseudocode"])); + assert_eq!(item["additionalProperties"], false); +} + +#[test] +fn each_provider_sends_its_own_request_and_reads_its_own_answer() { + let (_, sides) = project("a.py", "", LARGE); + let id = select(&trees(&sides), 3, None)[0].0; + let wrapped = format!("{{\"summaries\": {}}}", answers(&[(id, "call a, b, c")])); + // A compatible server that ignores the schema may wrap its answer. + let fenced = format!("maybe [a] or [b], or []\n```json\n{wrapped}\n```"); + for (provider, path, response, auth) in [ + ( + "openai", + "/v1", + json!({"choices": [{"message": {"role": "assistant", "content": fenced}}]}), + "authorization: bearer test-key", + ), + ( + "anthropic", + "", + json!({"content": [{"type": "thinking", "thinking": "[1]"}, {"type": "text", "text": wrapped}]}), + "x-api-key: test-key", + ), + ] { + let (file, mut sides) = project("a.py", "", LARGE); + let (endpoint, server) = serve_requests(vec![(200, response.to_string())]); + summarizer_with(json!({ + "provider": provider, + "model": "test-model", + "api_key": "test-key", + "endpoint": format!("{endpoint}{path}"), + "min_lines": 3, + "retries": 0, + })) + .run(&file, &mut sides) + .unwrap(); + assert_eq!(fold_label(&trees(&sides)), "call a, b, c", "{provider}"); + let request = server.join().unwrap().remove(0); + let body: serde_json::Value = serde_json::from_str(&request.body).unwrap(); + assert!( + request.headers.contains(&auth.to_owned()), + "{provider}: {:?}", + request.headers + ); + assert_eq!(body["model"], "test-model"); + match provider { + "openai" => { + assert_eq!(request.line, "POST /v1/chat/completions HTTP/1.1"); + assert_eq!(body["messages"][0]["role"], "system"); + assert!(body["messages"][1]["content"] + .as_str() + .unwrap() + .contains(&format!("fold {id}: lines 2-4"))); + assert!(body.get("temperature").is_none()); + let format = &body["response_format"]; + assert_eq!(format["type"], "json_schema"); + assert_eq!(format["json_schema"]["strict"], true); + assert_summaries_schema(&format["json_schema"]["schema"]); + } + _ => { + assert_eq!(request.line, "POST /v1/messages HTTP/1.1"); + assert!(request + .headers + .contains(&"anthropic-version: 2023-06-01".to_owned())); + assert_eq!(body["max_tokens"], 4096); + assert!(body.get("temperature").is_none()); + let format = &body["output_config"]["format"]; + assert_eq!(format["type"], "json_schema"); + assert_summaries_schema(&format["schema"]); + assert!(body["system"] + .as_str() + .unwrap() + .starts_with("For each listed fold")); + } + } + } +} + +#[test] +fn an_openai_compatible_server_needs_no_key() { + if std::env::var_os("OPENAI_API_KEY").is_some() { + return; + } + let (file, mut sides) = project("a.py", "", LARGE); + let id = select(&trees(&sides), 3, None)[0].0; + let response = json!({"choices": [{"message": {"content": answers(&[(id, "call a, b, c")])}}]}); + let (endpoint, server) = serve_requests(vec![(200, response.to_string())]); + summarizer_with(json!({ + "provider": "openai", + "model": "llama", + "endpoint": format!("{endpoint}/v1"), + "min_lines": 3, + })) + .run(&file, &mut sides) + .unwrap(); + assert_eq!(fold_label(&trees(&sides)), "call a, b, c"); + let request = server.join().unwrap().remove(0); + assert!(!request + .headers + .iter() + .any(|header| header.starts_with("authorization"))); +}