diff --git a/crates/switchyard-translation/src/lib.rs b/crates/switchyard-translation/src/lib.rs index d07a3497d..10da7ab08 100644 --- a/crates/switchyard-translation/src/lib.rs +++ b/crates/switchyard-translation/src/lib.rs @@ -31,5 +31,6 @@ pub use llm::*; pub use policy::*; pub use stream::*; pub use util::{ - PRESERVATION_METADATA_KEY, normalize_anthropic_tool_use_ids, sanitize_anthropic_tool_use_id, + PRESERVATION_METADATA_KEY, normalize_anthropic_tool_use_ids, prepare_request_for_target, + sanitize_anthropic_tool_use_id, }; diff --git a/crates/switchyard-translation/src/util.rs b/crates/switchyard-translation/src/util.rs index 9fcf769fd..11a5492d1 100644 --- a/crates/switchyard-translation/src/util.rs +++ b/crates/switchyard-translation/src/util.rs @@ -6,11 +6,12 @@ use std::collections::BTreeMap; use serde_json::{Map, Value, json}; +use switchyard_protocol::ModelId; use crate::diagnostic::TranslationDiagnostic; use crate::error::{Result, TranslationError}; use crate::format::FormatId; -use crate::llm::{ContentBlock, LlmRequest, Message, PreservationMetadata}; +use crate::llm::{ContentBlock, InstructionBlock, LlmRequest, Message, PreservationMetadata, Role}; use crate::policy::{ LossyConversionPolicy, PreservationPolicy, TranslationPolicy, UnknownFieldPolicy, }; @@ -271,6 +272,30 @@ pub fn exact_preserved_response( .flatten() } +/// Applies a selected target model and optionally prepends its system prompt. +/// +/// Adding a prompt invalidates preserved provider bodies because they predate the mutation. +/// Call this once per candidate using a request that has not already received a target prompt. +pub fn prepare_request_for_target( + request: &mut LlmRequest, + target: &ModelId, + prompt: Option<&str>, +) { + request.model = Some(target.to_string()); + if let Some(prompt) = prompt { + request.instructions.insert( + 0, + InstructionBlock { + role: Role::System, + content: vec![ContentBlock::Text { + text: prompt.to_string(), + }], + }, + ); + request.preservation.requests.clear(); + } +} + /// Embeds preservation metadata into a translated wire body when requested. pub fn embed_preservation( mut body: Value, diff --git a/crates/switchyard-translation/tests/request_translation.rs b/crates/switchyard-translation/tests/request_translation.rs index b7c4121a9..8310e7439 100644 --- a/crates/switchyard-translation/tests/request_translation.rs +++ b/crates/switchyard-translation/tests/request_translation.rs @@ -9,11 +9,69 @@ use pretty_assertions::assert_eq; use serde_json::{Value, json}; use switchyard_translation::{ LossyConversionPolicy, TranslationEngine, TranslationPolicy, WireFormat, + prepare_request_for_target, }; use common::{REASONING_MODEL, normalized_policy, shell_tool_call}; -type TestResult = std::result::Result<(), Box>; +type TestResult = std::result::Result>; + +// A target prompt makes every preserved provider body stale. +#[test] +fn preparing_a_target_prompt_invalidates_exact_replay() -> TestResult { + let engine = TranslationEngine::default(); + let policy = TranslationPolicy::default(); + let body = json!({ + "model": "route", + "messages": [ + {"role": "system", "name": "caller", "content": "client prompt"}, + {"role": "user", "content": "hi"} + ] + }); + let mut request = engine + .decode_request(WireFormat::OpenAiChat, &body, &policy)? + .request; + + prepare_request_for_target( + &mut request, + &"selected/model".into(), + Some("target prompt"), + ); + + assert!(request.preservation.requests.is_empty()); + let encoded = engine + .encode_request(WireFormat::OpenAiChat, &request, &policy)? + .body; + assert_eq!(encoded["model"], "selected/model"); + assert_eq!(encoded["messages"][0]["content"], "target prompt"); + assert_eq!(encoded["messages"][1]["content"], "client prompt"); + assert!(encoded["messages"][1].get("name").is_none()); + Ok(()) +} + +// Stamping only the normalized target does not invalidate exact replay. +#[test] +fn preparing_without_a_prompt_preserves_exact_replay() -> TestResult { + let engine = TranslationEngine::default(); + let policy = TranslationPolicy::default(); + let body = json!({ + "model": "route", + "messages": [{"role": "user", "content": "hi"}], + "provider_field": true + }); + let mut request = engine + .decode_request(WireFormat::OpenAiChat, &body, &policy)? + .request; + + prepare_request_for_target(&mut request, &"selected/model".into(), None); + + assert_eq!(request.model.as_deref(), Some("selected/model")); + assert_eq!( + request.preservation.requests[&WireFormat::OpenAiChat.into()], + body + ); + Ok(()) +} // Verifies Anthropic-only request fields are dropped or mapped for OpenAI Chat. #[test]