diff --git a/core/Cargo.lock b/core/Cargo.lock index de0d6b979..8a645583b 100644 --- a/core/Cargo.lock +++ b/core/Cargo.lock @@ -1656,6 +1656,14 @@ dependencies = [ "solstone-core-journal-io", ] +[[package]] +name = "solstone-core-generate" +version = "1.0.22" +dependencies = [ + "serde", + "serde_json", +] + [[package]] name = "solstone-core-indexer" version = "1.0.22" diff --git a/core/Cargo.toml b/core/Cargo.toml index 818505b9a..7ec244091 100644 --- a/core/Cargo.toml +++ b/core/Cargo.toml @@ -25,6 +25,7 @@ members = [ "crates/solstone-core-speakers", "crates/solstone-core-speakers-analyze", "crates/solstone-core-depict", + "crates/solstone-core-generate", "crates/solstone-core-speakers-onnx", ] resolver = "3" diff --git a/core/crates/solstone-core-generate/Cargo.toml b/core/crates/solstone-core-generate/Cargo.toml new file mode 100644 index 000000000..da447aafb --- /dev/null +++ b/core/crates/solstone-core-generate/Cargo.toml @@ -0,0 +1,16 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright (c) 2026 sol pbc + +[package] +name = "solstone-core-generate" +version.workspace = true +edition.workspace = true +rust-version.workspace = true +license.workspace = true + +[dependencies] +serde.workspace = true +serde_json.workspace = true + +[lints] +workspace = true diff --git a/core/crates/solstone-core-generate/src/client.rs b/core/crates/solstone-core-generate/src/client.rs new file mode 100644 index 000000000..3546278ab --- /dev/null +++ b/core/crates/solstone-core-generate/src/client.rs @@ -0,0 +1,93 @@ +// SPDX-License-Identifier: AGPL-3.0-only +// Copyright (c) 2026 sol pbc + +use std::collections::BTreeMap; +use std::env; +use std::io::Write; +use std::path::{Path, PathBuf}; +use std::process::{Command, Stdio}; + +use crate::{ + GenerateRequest, GenerateResponse, ProtocolError, decode_one_shot_response, + decode_protocol_error, encode_one_shot_request, +}; + +#[derive(Debug)] +pub enum ClientError { + Resolve(String), + Io(String), + Protocol(ProtocolError), + Decode(String), +} + +pub struct OneShotClient { + executable: PathBuf, + environment: BTreeMap, +} + +impl OneShotClient { + pub fn at_path(path: impl Into) -> Self { + Self { + executable: path.into(), + environment: BTreeMap::new(), + } + } + + pub fn with_env(mut self, name: impl Into, value: impl Into) -> Self { + self.environment.insert(name.into(), value.into()); + self + } + + pub fn sibling() -> Result { + let current = + env::current_exe().map_err(|error| ClientError::Resolve(error.to_string()))?; + let parent = current + .parent() + .ok_or_else(|| ClientError::Resolve("current executable has no parent".to_owned()))?; + let path = parent.join("solstone-generate-wire"); + if Path::new(&path).is_file() { + Ok(Self::at_path(path)) + } else { + Err(ClientError::Resolve(format!( + "missing sibling executable {}", + path.display() + ))) + } + } + + pub fn execute(&self, request: &GenerateRequest) -> Result { + let input = encode_one_shot_request(request).map_err(ClientError::Decode)?; + let mut child = Command::new(&self.executable) + .arg("--one-shot") + .envs(&self.environment) + .stdin(Stdio::piped()) + .stdout(Stdio::piped()) + .stderr(Stdio::piped()) + .spawn() + .map_err(|error| ClientError::Io(error.to_string()))?; + child + .stdin + .as_mut() + .ok_or_else(|| ClientError::Io("wire stdin is unavailable".to_owned()))? + .write_all(input.as_bytes()) + .map_err(|error| ClientError::Io(error.to_string()))?; + let output = child + .wait_with_output() + .map_err(|error| ClientError::Io(error.to_string()))?; + if output.status.success() { + decode_one_shot_response( + std::str::from_utf8(&output.stdout) + .map_err(|error| ClientError::Decode(error.to_string()))?, + ) + .map_err(ClientError::Decode) + } else { + Err(ClientError::Protocol( + decode_protocol_error( + std::str::from_utf8(&output.stderr) + .map_err(|error| ClientError::Decode(error.to_string()))?, + ) + .map_err(ClientError::Decode)?, + )) + } + } +} diff --git a/core/crates/solstone-core-generate/src/codec.rs b/core/crates/solstone-core-generate/src/codec.rs new file mode 100644 index 000000000..1beb07df5 --- /dev/null +++ b/core/crates/solstone-core-generate/src/codec.rs @@ -0,0 +1,481 @@ +// SPDX-License-Identifier: AGPL-3.0-only +// Copyright (c) 2026 sol pbc + +use std::collections::HashSet; + +use serde_json::{Map, Value, json}; + +use crate::fixture::schema; +use crate::types::{ + ContentPart, GenerateRequest, GenerateResponse, GeneratedResponse, ProtocolError, + ReasonCodeValue, RefusalReason, RefusedResponse, +}; + +fn object(value: Value) -> Result, String> { + value + .as_object() + .cloned() + .ok_or_else(|| "record must be a JSON object".to_owned()) +} + +fn string(object: &Map, name: &str) -> Result { + object + .get(name) + .and_then(Value::as_str) + .map(ToOwned::to_owned) + .ok_or_else(|| format!("{name} must be a string")) +} + +fn optional_string(object: &Map, name: &str) -> Result, String> { + match object.get(name) { + None | Some(Value::Null) => Ok(None), + Some(Value::String(value)) => Ok(Some(value.clone())), + _ => Err(format!("{name} must be a string or null")), + } +} + +fn optional_value(object: &Map, name: &str) -> Option { + object.get(name).filter(|value| !value.is_null()).cloned() +} + +fn require_schema(object: &Map, expected: &str) -> Result<(), String> { + if object.get("schema").and_then(Value::as_str) == Some(expected) { + Ok(()) + } else { + Err("record schema is not supported".to_owned()) + } +} + +pub fn encode_one_shot_request(request: &GenerateRequest) -> Result { + if request.contents.is_empty() { + return Err("contents must be non-empty".to_owned()); + } + let contents = request + .contents + .iter() + .map(|part| match part { + ContentPart::Text { text } => json!({"type": "text", "text": text}), + ContentPart::Image { mime_type, data } => { + json!({"type": "image", "mime_type": mime_type, "data": data}) + } + }) + .collect::>(); + serde_json::to_string(&json!({ + "schema": schema("request"), + "id": request.id, + "context": request.context, + "contents": contents, + "system_instruction": request.system_instruction, + "temperature": request.temperature, + "max_output_tokens": request.max_output_tokens, + "thinking_budget": request.thinking_budget, + "timeout_s": request.timeout_s, + "json_output": request.json_output, + "json_schema": request.json_schema, + "enforce_responsiveness": request.enforce_responsiveness, + "attempt_index": request.attempt_index, + "exclusive_admission": request.exclusive_admission, + "transport_retries": request.transport_retries, + })) + .map_err(|error| error.to_string()) +} + +pub fn decode_one_shot_request(input: &str) -> Result { + let object = object(serde_json::from_str(input).map_err(|error| error.to_string())?)?; + require_schema(&object, schema("request"))?; + let contents = object + .get("contents") + .and_then(Value::as_array) + .ok_or_else(|| "contents must be an array".to_owned())? + .iter() + .map(|value| { + let part = value + .as_object() + .ok_or_else(|| "content must be an object".to_owned())?; + match part.get("type").and_then(Value::as_str) { + Some("text") => Ok(ContentPart::Text { + text: string(part, "text")?, + }), + Some("image") => Ok(ContentPart::Image { + mime_type: string(part, "mime_type")?, + data: string(part, "data")?, + }), + _ => Err("content type is not supported".to_owned()), + } + }) + .collect::, _>>()?; + if contents.is_empty() { + return Err("contents must be non-empty".to_owned()); + } + Ok(GenerateRequest { + id: optional_string(&object, "id")?, + context: string(&object, "context")?, + contents, + system_instruction: optional_string(&object, "system_instruction")?, + temperature: object + .get("temperature") + .and_then(Value::as_f64) + .ok_or_else(|| "temperature must be a number".to_owned())?, + max_output_tokens: object + .get("max_output_tokens") + .and_then(Value::as_u64) + .ok_or_else(|| "max_output_tokens must be an integer".to_owned())?, + thinking_budget: object.get("thinking_budget").and_then(Value::as_u64), + timeout_s: object.get("timeout_s").and_then(Value::as_f64), + json_output: object + .get("json_output") + .and_then(Value::as_bool) + .ok_or_else(|| "json_output must be a boolean".to_owned())?, + json_schema: optional_value(&object, "json_schema"), + enforce_responsiveness: object + .get("enforce_responsiveness") + .and_then(Value::as_bool) + .ok_or_else(|| "enforce_responsiveness must be a boolean".to_owned())?, + attempt_index: object + .get("attempt_index") + .and_then(Value::as_u64) + .ok_or_else(|| "attempt_index must be an integer".to_owned())?, + exclusive_admission: object + .get("exclusive_admission") + .and_then(Value::as_bool) + .ok_or_else(|| "exclusive_admission must be a boolean".to_owned())?, + transport_retries: object.get("transport_retries").and_then(Value::as_u64), + }) +} + +pub fn decode_one_shot_response(input: &str) -> Result { + let object = object(serde_json::from_str(input).map_err(|error| error.to_string())?)?; + require_schema(&object, schema("response"))?; + match string(&object, "outcome")?.as_str() { + "generated" => Ok(GenerateResponse::Generated(Box::new(GeneratedResponse { + id: optional_string(&object, "id")?, + text: string(&object, "text")?, + model: string(&object, "model")?, + usage: object + .get("usage") + .cloned() + .ok_or_else(|| "usage is required".to_owned())?, + finish_reason: string(&object, "finish_reason")?, + thinking: optional_value(&object, "thinking"), + schema_validation: optional_value(&object, "schema_validation"), + input_budget: optional_value(&object, "input_budget"), + request_budget: optional_value(&object, "request_budget"), + inference: optional_value(&object, "inference"), + }))), + "refused" => { + let reason_code = + optional_string(&object, "reason_code")?.map(ReasonCodeValue::from_wire); + let mut retryable = object + .get("retryable") + .and_then(Value::as_bool) + .ok_or_else(|| "retryable must be a boolean".to_owned())?; + let mut blocking = object + .get("blocking") + .and_then(Value::as_bool) + .ok_or_else(|| "blocking must be a boolean".to_owned())?; + if matches!(reason_code, Some(ReasonCodeValue::Unknown(_))) { + retryable = false; + blocking = true; + } + Ok(GenerateResponse::Refused(RefusedResponse { + id: optional_string(&object, "id")?, + reason: RefusalReason::from_wire(&string(&object, "reason")?), + reason_code, + retryable, + blocking, + reset_at_ms: object.get("reset_at_ms").and_then(Value::as_u64), + provider: optional_string(&object, "provider")?, + detail: string(&object, "detail")?, + })) + } + _ => Err("response outcome is not supported".to_owned()), + } +} + +pub fn decode_protocol_error(input: &str) -> Result { + let object = object(serde_json::from_str(input).map_err(|error| error.to_string())?)?; + require_schema(&object, schema("error"))?; + Ok(ProtocolError { + id: optional_string(&object, "id")?, + reason: string(&object, "reason")?, + detail: string(&object, "detail")?, + }) +} + +pub fn encode_session_request_line(request: &GenerateRequest) -> Result { + if request.id.is_none() { + return Err("session request id is required".to_owned()); + } + Ok(format!("{}\n", encode_one_shot_request(request)?)) +} + +pub fn decode_session_request_line(line: &str) -> Result { + let request = decode_one_shot_request(line.trim_end())?; + if request.id.is_none() { + return Err("session request id is required".to_owned()); + } + Ok(request) +} + +pub fn decode_session_response_line(line: &str) -> Result { + let response = decode_one_shot_response(line.trim_end())?; + let id = match &response { + GenerateResponse::Generated(value) => &value.id, + GenerateResponse::Refused(value) => &value.id, + }; + if id.is_none() { + return Err("session response id is required".to_owned()); + } + Ok(response) +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum SessionError { + MissingId, + UnknownOrRetiredId(String), + Terminal, +} + +#[derive(Debug, Default)] +pub struct SessionCorrelation { + outstanding: HashSet, + retired: HashSet, + terminal: bool, +} + +impl SessionCorrelation { + pub fn submit(&mut self, id: impl Into) -> Result<(), SessionError> { + if self.terminal { + return Err(SessionError::Terminal); + } + let id = id.into(); + if id.is_empty() { + return Err(SessionError::MissingId); + } + self.outstanding.insert(id); + Ok(()) + } + + pub fn accept(&mut self, response: &GenerateResponse) -> Result<(), SessionError> { + if self.terminal { + return Err(SessionError::Terminal); + } + let id = match response { + GenerateResponse::Generated(value) => value.id.as_deref(), + GenerateResponse::Refused(value) => value.id.as_deref(), + } + .ok_or(SessionError::MissingId)?; + if !self.outstanding.remove(id) || self.retired.contains(id) { + self.terminal = true; + return Err(SessionError::UnknownOrRetiredId(id.to_owned())); + } + self.retired.insert(id.to_owned()); + Ok(()) + } +} + +#[cfg(test)] +mod tests { + use serde_json::json; + + use super::*; + use crate::contract; + use crate::types::{Outcome, ReasonCodeValue}; + + fn request(id: Option<&str>) -> GenerateRequest { + GenerateRequest { + id: id.map(ToOwned::to_owned), + context: "test.generate".to_owned(), + contents: vec![ContentPart::Text { + text: "OK".to_owned(), + }], + system_instruction: None, + temperature: 0.3, + max_output_tokens: 16, + thinking_budget: None, + timeout_s: None, + json_output: false, + json_schema: None, + enforce_responsiveness: true, + attempt_index: 0, + exclusive_admission: false, + transport_retries: None, + } + } + + #[test] + fn fixture_conformance_vectors_decode() { + for vector in contract()["conformance_vectors"].as_array().unwrap() { + match vector["framing"].as_str().unwrap() { + "one_shot" => { + if let Some(response) = vector.get("response") { + decode_one_shot_response(&response.to_string()).unwrap(); + } + if let Some(request) = vector.get("request") { + decode_one_shot_request(&request.to_string()).unwrap(); + } + } + "protocol_error" => { + decode_protocol_error(&vector["protocol_error"].to_string()).unwrap(); + } + framing => panic!("unexpected fixture framing {framing}"), + } + } + } + + #[test] + fn enum_members_match_fixture() { + let outcomes = [Outcome::Generated, Outcome::Refused] + .into_iter() + .map(Outcome::as_str) + .collect::>(); + assert_eq!( + outcomes, + contract()["outcomes"] + .as_array() + .unwrap() + .iter() + .map(|value| value.as_str().unwrap()) + .collect::>() + ); + let reasons = [ + RefusalReason::AttestationNotVerified, + RefusalReason::AttestationFailed, + RefusalReason::AttestationStale, + RefusalReason::NoEngineConfigured, + RefusalReason::IncompleteJson, + RefusalReason::IncompleteText, + RefusalReason::ProviderResponseInvalid, + RefusalReason::SchemaValidationFailed, + RefusalReason::NonResponsiveOutput, + RefusalReason::Unknown, + ] + .into_iter() + .map(RefusalReason::as_str) + .collect::>(); + assert_eq!( + reasons, + contract()["refusal_reasons"] + .as_array() + .unwrap() + .iter() + .map(|value| value.as_str().unwrap()) + .collect::>() + ); + } + + #[test] + fn tagged_union_round_trips_legal_variants_only() { + let generated = + decode_one_shot_response(&contract()["conformance_vectors"][0]["response"].to_string()) + .unwrap(); + let refused = + decode_one_shot_response(&contract()["conformance_vectors"][1]["response"].to_string()) + .unwrap(); + assert!(matches!(generated, GenerateResponse::Generated(_))); + assert!(matches!(refused, GenerateResponse::Refused(_))); + } + + #[test] + fn refused_round_trip_preserves_classification_and_metadata() { + let vector = &contract()["conformance_vectors"][1]["response"]; + let response = decode_one_shot_response(&vector.to_string()).unwrap(); + let GenerateResponse::Refused(value) = response else { + panic!("expected refusal") + }; + assert_eq!(value.reason, RefusalReason::AttestationNotVerified); + assert_eq!( + value.reason_code.unwrap().as_wire(), + "attestation_not_yet_verified" + ); + assert!(!value.retryable || value.blocking); + assert_eq!(value.reset_at_ms, None); + assert_eq!(value.provider.as_deref(), Some("local")); + } + + #[test] + fn generated_round_trip_preserves_all_result_fields() { + let response = + decode_one_shot_response(&contract()["conformance_vectors"][0]["response"].to_string()) + .unwrap(); + let GenerateResponse::Generated(value) = response else { + panic!("expected generated") + }; + assert_eq!(value.text, "OK"); + assert_eq!(value.model, "fixture-model"); + assert_eq!( + value.usage, + json!({"input_tokens": 2, "output_tokens": 1, "total_tokens": 3}) + ); + assert_eq!(value.finish_reason, "stop"); + assert_eq!(value.thinking, None); + assert_eq!(value.schema_validation, None); + assert_eq!(value.input_budget, None); + assert_eq!(value.request_budget, None); + assert_eq!(value.inference, None); + } + + #[test] + fn unknown_reason_code_is_safe() { + let vector = &contract()["conformance_vectors"][12]["response"]; + let response = decode_one_shot_response(&vector.to_string()).unwrap(); + let GenerateResponse::Refused(value) = response else { + panic!("expected refusal") + }; + assert!(matches!( + value.reason_code, + Some(ReasonCodeValue::Unknown(_)) + )); + assert!(!value.retryable); + assert!(value.blocking); + } + + #[test] + fn wrong_schema_is_rejected() { + let mut value = contract()["conformance_vectors"][0]["response"].clone(); + value["schema"] = json!("wrong"); + assert!(decode_one_shot_response(&value.to_string()).is_err()); + } + + #[test] + fn session_codec_correlates_out_of_order_ids_and_rejects_missing_ids() { + let first = request(Some("first")); + let second = request(Some("second")); + assert!(encode_session_request_line(&first).unwrap().ends_with('\n')); + assert!( + decode_session_request_line(&encode_session_request_line(&second).unwrap()).is_ok() + ); + assert!(encode_session_request_line(&request(None)).is_err()); + let mut tracker = SessionCorrelation::default(); + tracker.submit("first").unwrap(); + tracker.submit("second").unwrap(); + let first_response = GenerateResponse::Generated(Box::new(GeneratedResponse { + id: Some("first".to_owned()), + text: "one".to_owned(), + model: "m".to_owned(), + usage: json!({}), + finish_reason: "stop".to_owned(), + thinking: None, + schema_validation: None, + input_budget: None, + request_budget: None, + inference: None, + })); + let second_response = GenerateResponse::Refused(RefusedResponse { + id: Some("second".to_owned()), + reason: RefusalReason::NoEngineConfigured, + reason_code: None, + retryable: false, + blocking: true, + reset_at_ms: None, + provider: Some("none".to_owned()), + detail: "none".to_owned(), + }); + tracker.accept(&second_response).unwrap(); + tracker.accept(&first_response).unwrap(); + assert_eq!( + tracker.accept(&first_response), + Err(SessionError::UnknownOrRetiredId("first".to_owned())) + ); + } +} diff --git a/core/crates/solstone-core-generate/src/fixture.rs b/core/crates/solstone-core-generate/src/fixture.rs new file mode 100644 index 000000000..df6c379df --- /dev/null +++ b/core/crates/solstone-core-generate/src/fixture.rs @@ -0,0 +1,29 @@ +// SPDX-License-Identifier: AGPL-3.0-only +// Copyright (c) 2026 sol pbc + +use std::sync::OnceLock; + +use serde_json::Value; + +const CONTRACT_FIXTURE: &str = include_str!("../../../fixtures/generate_contract.json"); +static CONTRACT: OnceLock = OnceLock::new(); + +pub fn contract() -> &'static Value { + CONTRACT.get_or_init(|| { + serde_json::from_str(CONTRACT_FIXTURE).expect("generate contract fixture is valid JSON") + }) +} + +pub(crate) fn schema(name: &str) -> &'static str { + contract()["schema_identifiers"][name] + .as_str() + .expect("generate contract schema identifier is a string") +} + +pub(crate) fn known_reason_code(code: &str) -> bool { + contract()["reason_codes"] + .as_array() + .expect("generate contract reason codes are an array") + .iter() + .any(|entry| entry["code"].as_str() == Some(code)) +} diff --git a/core/crates/solstone-core-generate/src/lib.rs b/core/crates/solstone-core-generate/src/lib.rs new file mode 100644 index 000000000..2b73fb9ed --- /dev/null +++ b/core/crates/solstone-core-generate/src/lib.rs @@ -0,0 +1,21 @@ +// SPDX-License-Identifier: AGPL-3.0-only +// Copyright (c) 2026 sol pbc + +//! Typed codecs and one-shot client for the generate contract. + +mod client; +mod codec; +mod fixture; +mod types; + +pub use client::{ClientError, OneShotClient}; +pub use codec::{ + SessionCorrelation, SessionError, decode_one_shot_request, decode_one_shot_response, + decode_protocol_error, decode_session_request_line, decode_session_response_line, + encode_one_shot_request, encode_session_request_line, +}; +pub use fixture::contract; +pub use types::{ + ContentPart, GenerateRequest, GenerateResponse, GeneratedResponse, Outcome, ProtocolError, + ReasonCode, ReasonCodeValue, RefusalReason, RefusedResponse, UnknownReasonCode, +}; diff --git a/core/crates/solstone-core-generate/src/types.rs b/core/crates/solstone-core-generate/src/types.rs new file mode 100644 index 000000000..c11816106 --- /dev/null +++ b/core/crates/solstone-core-generate/src/types.rs @@ -0,0 +1,179 @@ +// SPDX-License-Identifier: AGPL-3.0-only +// Copyright (c) 2026 sol pbc + +use serde_json::Value; + +use crate::fixture::known_reason_code; + +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum ContentPart { + Text { text: String }, + Image { mime_type: String, data: String }, +} + +#[derive(Debug, Clone, PartialEq)] +pub struct GenerateRequest { + pub id: Option, + pub context: String, + pub contents: Vec, + pub system_instruction: Option, + pub temperature: f64, + pub max_output_tokens: u64, + pub thinking_budget: Option, + pub timeout_s: Option, + pub json_output: bool, + pub json_schema: Option, + pub enforce_responsiveness: bool, + pub attempt_index: u64, + pub exclusive_admission: bool, + pub transport_retries: Option, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum Outcome { + Generated, + Refused, +} + +impl Outcome { + pub const fn as_str(self) -> &'static str { + match self { + Self::Generated => "generated", + Self::Refused => "refused", + } + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum RefusalReason { + AttestationNotVerified, + AttestationFailed, + AttestationStale, + NoEngineConfigured, + IncompleteJson, + IncompleteText, + ProviderResponseInvalid, + SchemaValidationFailed, + NonResponsiveOutput, + Unknown, +} + +impl RefusalReason { + pub const fn as_str(self) -> &'static str { + match self { + Self::AttestationNotVerified => "attestation-not-verified", + Self::AttestationFailed => "attestation-failed", + Self::AttestationStale => "attestation-stale", + Self::NoEngineConfigured => "no-engine-configured", + Self::IncompleteJson => "incomplete-json", + Self::IncompleteText => "incomplete-text", + Self::ProviderResponseInvalid => "provider-response-invalid", + Self::SchemaValidationFailed => "schema-validation-failed", + Self::NonResponsiveOutput => "non-responsive-output", + Self::Unknown => "unknown", + } + } + + pub(crate) fn from_wire(value: &str) -> Self { + match value { + "attestation-not-verified" => Self::AttestationNotVerified, + "attestation-failed" => Self::AttestationFailed, + "attestation-stale" => Self::AttestationStale, + "no-engine-configured" => Self::NoEngineConfigured, + "incomplete-json" => Self::IncompleteJson, + "incomplete-text" => Self::IncompleteText, + "provider-response-invalid" => Self::ProviderResponseInvalid, + "schema-validation-failed" => Self::SchemaValidationFailed, + "non-responsive-output" => Self::NonResponsiveOutput, + _ => Self::Unknown, + } + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct ReasonCode(String); + +impl ReasonCode { + pub fn new(value: impl Into) -> Result { + let value = value.into(); + if known_reason_code(&value) { + Ok(Self(value)) + } else { + Err(format!("unknown generate reason code {value:?}")) + } + } + + pub fn as_str(&self) -> &str { + &self.0 + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct UnknownReasonCode { + pub received: String, + pub canonical: ReasonCode, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum ReasonCodeValue { + Known(ReasonCode), + Unknown(UnknownReasonCode), +} + +impl ReasonCodeValue { + pub(crate) fn from_wire(value: String) -> Self { + match ReasonCode::new(value.clone()) { + Ok(code) => Self::Known(code), + Err(_) => Self::Unknown(UnknownReasonCode { + received: value, + canonical: ReasonCode::new("unknown").expect("unknown code is in fixture"), + }), + } + } + + pub fn as_wire(&self) -> &str { + match self { + Self::Known(code) => code.as_str(), + Self::Unknown(code) => &code.received, + } + } +} + +#[derive(Debug, Clone, PartialEq)] +pub struct GeneratedResponse { + pub id: Option, + pub text: String, + pub model: String, + pub usage: Value, + pub finish_reason: String, + pub thinking: Option, + pub schema_validation: Option, + pub input_budget: Option, + pub request_budget: Option, + pub inference: Option, +} + +#[derive(Debug, Clone, PartialEq)] +pub struct RefusedResponse { + pub id: Option, + pub reason: RefusalReason, + pub reason_code: Option, + pub retryable: bool, + pub blocking: bool, + pub reset_at_ms: Option, + pub provider: Option, + pub detail: String, +} + +#[derive(Debug, Clone, PartialEq)] +pub enum GenerateResponse { + Generated(Box), + Refused(RefusedResponse), +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct ProtocolError { + pub id: Option, + pub reason: String, + pub detail: String, +} diff --git a/core/crates/solstone-core-generate/tests/support.rs b/core/crates/solstone-core-generate/tests/support.rs new file mode 100644 index 000000000..b90547c27 --- /dev/null +++ b/core/crates/solstone-core-generate/tests/support.rs @@ -0,0 +1,33 @@ +// SPDX-License-Identifier: AGPL-3.0-only +// Copyright (c) 2026 sol pbc + +use std::env; +use std::path::PathBuf; + +pub fn generate_wire() -> PathBuf { + if let Some(path) = env::var_os("SOLSTONE_GENERATE_WIRE") { + let path = PathBuf::from(path); + assert!( + path.is_file(), + "SOLSTONE_GENERATE_WIRE is not a file: {}", + path.display() + ); + return path; + } + let manifest = PathBuf::from(env!("CARGO_MANIFEST_DIR")); + let root = manifest + .ancestors() + .find(|candidate| { + candidate.join(".git").exists() + && candidate.join("pyproject.toml").is_file() + && candidate.join("solstone").is_dir() + }) + .expect("generate-wire integration tests require a checkout root"); + let path = root.join(".venv/bin/solstone-generate-wire"); + assert!( + path.is_file(), + "missing {}; set SOLSTONE_GENERATE_WIRE or run make install", + path.display() + ); + path +} diff --git a/core/crates/solstone-core-generate/tests/wire.rs b/core/crates/solstone-core-generate/tests/wire.rs new file mode 100644 index 000000000..97c77d705 --- /dev/null +++ b/core/crates/solstone-core-generate/tests/wire.rs @@ -0,0 +1,400 @@ +// SPDX-License-Identifier: AGPL-3.0-only +// Copyright (c) 2026 sol pbc + +mod support; + +use std::fs; +use std::io::{Read, Write}; +use std::net::{TcpListener, TcpStream}; +use std::path::PathBuf; +use std::process::{Command, Stdio}; +use std::thread; +use std::time::{SystemTime, UNIX_EPOCH}; + +use solstone_core_generate::{ + ContentPart, GenerateRequest, GenerateResponse, OneShotClient, RefusalReason, contract, + encode_one_shot_request, +}; + +struct Journal { + path: PathBuf, +} + +impl Journal { + fn no_engine() -> Self { + let path = std::env::temp_dir().join(format!( + "solstone-generate-test-{}-{}", + std::process::id(), + SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap() + .as_nanos(), + )); + fs::create_dir_all(path.join("config")).unwrap(); + fs::write( + path.join("config/journal.json"), + r#"{"providers":{"active":{"provider":"none"}}}"#, + ) + .unwrap(); + Self { path } + } + + fn bundled_local(port: u16) -> Self { + let journal = Self::no_engine(); + fs::write( + journal.path.join("config/journal.json"), + r#"{"providers":{"active":{"provider":"local"}}}"#, + ) + .unwrap(); + fs::create_dir_all(journal.path.join("health")).unwrap(); + fs::write(journal.path.join("health/local.port"), port.to_string()).unwrap(); + journal + } + + fn byo_unreachable() -> Self { + let journal = Self::no_engine(); + fs::write( + journal.path.join("config/journal.json"), + r#"{"providers":{"active":{"provider":"local"},"local":{"endpoint_url":"http://127.0.0.1:1","served_model_id":"stub"}}}"#, + ) + .unwrap(); + journal + } +} + +impl Drop for Journal { + fn drop(&mut self) { + let _ = fs::remove_dir_all(&self.path); + } +} + +fn request() -> GenerateRequest { + GenerateRequest { + id: Some("wire-test".to_owned()), + context: "test.generate".to_owned(), + contents: vec![ContentPart::Text { + text: "OK".to_owned(), + }], + system_instruction: None, + temperature: 0.3, + max_output_tokens: 16, + thinking_budget: None, + timeout_s: Some(3.0), + json_output: false, + json_schema: None, + enforce_responsiveness: true, + attempt_index: 0, + exclusive_admission: false, + transport_retries: None, + } +} + +struct LocalStub { + port: u16, + worker: thread::JoinHandle<()>, +} + +#[derive(Clone, Copy)] +struct Completion { + text: &'static str, + finish_reason: &'static str, +} + +impl LocalStub { + fn start() -> Self { + Self::with_completion(Completion { + text: "OK", + finish_reason: "stop", + }) + } + + fn with_completion(completion: Completion) -> Self { + let listener = TcpListener::bind("127.0.0.1:0").unwrap(); + let port = listener.local_addr().unwrap().port(); + let worker = thread::spawn(move || { + for stream in listener.incoming().take(12) { + let completed = handle_local_request(stream.unwrap(), completion); + if completed { + return; + } + } + }); + Self { port, worker } + } + + fn finish(self) { + self.worker.join().unwrap(); + } +} + +fn handle_local_request(mut stream: TcpStream, completion: Completion) -> bool { + let mut request = Vec::new(); + let mut chunk = [0_u8; 4096]; + loop { + let read = stream.read(&mut chunk).unwrap(); + request.extend_from_slice(&chunk[..read]); + let Some(header_end) = request.windows(4).position(|window| window == b"\r\n\r\n") else { + continue; + }; + let header = String::from_utf8_lossy(&request[..header_end]); + let content_length = header + .lines() + .find_map(|line| line.strip_prefix("Content-Length: ")) + .and_then(|value| value.parse::().ok()) + .unwrap_or_default(); + if request.len() >= header_end + 4 + content_length { + break; + } + } + let head = String::from_utf8_lossy(&request); + let request_line = head.lines().next().unwrap_or_default(); + let (body, completed) = if request_line.starts_with("GET /health ") { + (r#"{"loaded_model":"local"}"#.to_owned(), false) + } else if request_line.starts_with("GET /props ") { + (r#"{"n_ctx":16384,"total_slots":1}"#.to_owned(), false) + } else if request_line.starts_with("POST /tokenize ") { + (r#"{"tokens":[1]}"#.to_owned(), false) + } else if request_line.starts_with("POST /v1/chat/completions ") { + ( + format!( + r#"{{"choices":[{{"message":{{"content":{}}},"finish_reason":{}}}],"usage":{{"prompt_tokens":1,"completion_tokens":1,"total_tokens":2}}}}"#, + serde_json::to_string(completion.text).unwrap(), + serde_json::to_string(completion.finish_reason).unwrap(), + ), + true, + ) + } else { + ("{}".to_owned(), false) + }; + write!( + stream, + "HTTP/1.1 200 OK\r\nContent-Type: application/json\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{}", + body.len(), + body + ) + .unwrap(); + completed +} + +fn spawn_v2(journal: &Journal, request: &GenerateRequest) -> std::process::Output { + Command::new(support::generate_wire()) + .arg("--one-shot") + .env("SOLSTONE_JOURNAL", &journal.path) + .stdin(Stdio::piped()) + .stdout(Stdio::piped()) + .stderr(Stdio::piped()) + .spawn() + .and_then(|mut child| { + child + .stdin + .as_mut() + .unwrap() + .write_all(encode_one_shot_request(request).unwrap().as_bytes())?; + child.wait_with_output() + }) + .unwrap() +} + +fn bundled_refusal( + completion: Completion, + configure: impl FnOnce(&mut GenerateRequest), + reason: RefusalReason, + reason_code: Option<&str>, +) { + let stub = LocalStub::with_completion(completion); + let journal = Journal::bundled_local(stub.port); + let mut generated_request = request(); + configure(&mut generated_request); + let output = spawn_v2(&journal, &generated_request); + stub.finish(); + assert_eq!(output.status.code(), Some(0)); + assert!(matches!(output.status.code(), Some(0 | 64 | 70))); + assert_ne!(output.status.code(), Some(69)); + let response = solstone_core_generate::decode_one_shot_response( + std::str::from_utf8(&output.stdout).unwrap(), + ) + .unwrap(); + let GenerateResponse::Refused(refusal) = response else { + panic!("expected refused response") + }; + assert_eq!(refusal.reason, reason); + assert_eq!( + refusal.reason_code.as_ref().map(|code| code.as_wire()), + reason_code + ); +} + +#[test] +fn one_shot_client_round_trips_no_engine_refusal() { + let journal = Journal::no_engine(); + let response = OneShotClient::at_path(support::generate_wire()) + .with_env("SOLSTONE_JOURNAL", journal.path.to_string_lossy()) + .execute(&request()) + .unwrap(); + let GenerateResponse::Refused(refusal) = response else { + panic!("expected refusal") + }; + assert_eq!(refusal.reason, RefusalReason::NoEngineConfigured); + assert_eq!(refusal.provider.as_deref(), Some("none")); +} + +#[test] +fn v2_request_has_no_provider_or_model_field() { + let value: serde_json::Value = + serde_json::from_str(&encode_one_shot_request(&request()).unwrap()).unwrap(); + assert!(value.get("provider").is_none()); + assert!(value.get("model").is_none()); +} + +#[test] +fn real_wire_rejects_unknown_v2_request_field_on_stderr() { + let journal = Journal::no_engine(); + let mut value: serde_json::Value = + serde_json::from_str(&encode_one_shot_request(&request()).unwrap()).unwrap(); + value["unknown"] = serde_json::json!(true); + let output = Command::new(support::generate_wire()) + .arg("--one-shot") + .env("SOLSTONE_JOURNAL", &journal.path) + .stdin(Stdio::piped()) + .stdout(Stdio::piped()) + .stderr(Stdio::piped()) + .spawn() + .and_then(|mut child| { + use std::io::Write; + child + .stdin + .as_mut() + .unwrap() + .write_all(value.to_string().as_bytes())?; + child.wait_with_output() + }) + .unwrap(); + assert_eq!(output.status.code(), Some(64)); + assert!(matches!(output.status.code(), Some(0 | 64 | 70))); + assert_ne!(output.status.code(), Some(69)); + assert!(output.stdout.is_empty()); + let error: serde_json::Value = serde_json::from_slice(&output.stderr).unwrap(); + assert_eq!(error["schema"], contract()["schema_identifiers"]["error"]); +} + +#[test] +fn real_wire_contract_matches_compiled_fixture_bytes() { + let output = Command::new(support::generate_wire()) + .arg("--contract") + .output() + .unwrap(); + assert!(output.status.success()); + let expected = serde_json::to_vec_pretty(contract()).unwrap(); + assert_eq!(output.stdout, [expected, b"\n".to_vec()].concat()); +} + +#[test] +fn bundled_local_round_trip_generates_and_logs_one_usage_record() { + let stub = LocalStub::start(); + let journal = Journal::bundled_local(stub.port); + let output = spawn_v2(&journal, &request()); + stub.finish(); + assert_eq!(output.status.code(), Some(0)); + assert!(String::from_utf8_lossy(&output.stderr).contains("provider=local")); + assert_eq!( + output.stdout.iter().filter(|byte| **byte == b'\n').count(), + 1 + ); + let response = solstone_core_generate::decode_one_shot_response( + std::str::from_utf8(&output.stdout).unwrap(), + ) + .unwrap(); + let GenerateResponse::Generated(generated) = response else { + panic!("expected generated response: {response:?}") + }; + assert_eq!(generated.text, "OK"); + let token_lines = fs::read_dir(journal.path.join("tokens")) + .unwrap() + .map(|entry| { + fs::read_to_string(entry.unwrap().path()) + .unwrap() + .lines() + .count() + }) + .sum::(); + assert_eq!(token_lines, 1); +} + +#[test] +fn bundled_local_incomplete_json_refuses() { + bundled_refusal( + Completion { + text: "{}", + finish_reason: "length", + }, + |request| request.json_output = true, + RefusalReason::IncompleteJson, + Some("incomplete_json_length"), + ); +} + +#[test] +fn bundled_local_invalid_provider_response_refuses() { + bundled_refusal( + Completion { + text: "", + finish_reason: "stop", + }, + |_| {}, + RefusalReason::ProviderResponseInvalid, + Some("provider_response_invalid"), + ); +} + +#[test] +fn bundled_local_non_responsive_output_refuses() { + bundled_refusal( + Completion { + text: "I cannot help with that request.", + finish_reason: "stop", + }, + |_| {}, + RefusalReason::NonResponsiveOutput, + Some("non_responsive"), + ); +} + +#[test] +fn byo_unreachable_refuses_with_diagnostics_and_one_stdout_record() { + let journal = Journal::byo_unreachable(); + let output = Command::new(support::generate_wire()) + .arg("--one-shot") + .env("SOLSTONE_JOURNAL", &journal.path) + .stdin(Stdio::piped()) + .stdout(Stdio::piped()) + .stderr(Stdio::piped()) + .spawn() + .and_then(|mut child| { + child + .stdin + .as_mut() + .unwrap() + .write_all(encode_one_shot_request(&request()).unwrap().as_bytes())?; + child.wait_with_output() + }) + .unwrap(); + assert_eq!(output.status.code(), Some(0)); + assert!(matches!(output.status.code(), Some(0 | 64 | 70))); + assert_ne!(output.status.code(), Some(69)); + let diagnostics = String::from_utf8_lossy(&output.stderr); + assert!(diagnostics.contains("provider=local")); + assert!( + diagnostics.contains("solstone-generate-wire v2:"), + "expected wire diagnostic, got: {diagnostics}" + ); + assert!(diagnostics.lines().count() >= 2); + assert_eq!( + output.stdout.iter().filter(|byte| **byte == b'\n').count(), + 1 + ); + let response = solstone_core_generate::decode_one_shot_response( + std::str::from_utf8(&output.stdout).unwrap(), + ) + .unwrap(); + assert!(matches!(response, GenerateResponse::Refused(_))); +}