1use std::borrow::Cow;
8
9use serde::Serialize;
10use serde::de::DeserializeOwned;
11use serde_json::{Map, Value};
12
13use super::mapping::Mapping;
14use super::{CacheRetention, GenerationOptions, OnUnsupported, UnsupportedOption};
15use crate::completion::provider_options::SHARED;
16use crate::completion::{CompletionRequest, ReplayTarget};
17use crate::error::EncodeError;
18use crate::providers::openai::wire::BodyRewrite;
19use crate::wire::Body;
20
21#[derive(Clone, Debug, PartialEq, Serialize)]
25#[serde(transparent)]
26pub struct FinalBody(Map<String, Value>);
27
28impl FinalBody {
29 pub fn is_empty(&self) -> bool {
31 self.0.is_empty()
32 }
33
34 pub fn get(&self, key: &str) -> Option<&Value> {
36 self.0.get(key)
37 }
38
39 pub fn pointer(&self, pointer: &str) -> Option<&Value> {
41 let path = pointer.strip_prefix('/')?;
42 let (head, rest) = path.split_once('/').unwrap_or((path, ""));
43 let value = self.0.get(&head.replace("~1", "/").replace("~0", "~"))?;
44 if rest.is_empty() {
45 Some(value)
46 } else {
47 value.pointer(&format!("/{rest}"))
48 }
49 }
50
51 pub fn into_body(self) -> Body {
54 Body::Bytes(Value::Object(self.0).to_string().into_bytes())
55 }
56
57 pub fn deserialize<T: DeserializeOwned>(&self) -> Result<T, serde_json::Error> {
64 T::deserialize(&Value::Object(self.0.clone()))
65 }
66}
67
68#[derive(Clone, Copy, Debug, PartialEq, Eq)]
70#[non_exhaustive]
71pub enum RawAt {
72 Top,
74 Under(&'static str),
77 Split {
80 top: &'static [&'static str],
82 rest: &'static str,
84 },
85 #[doc(hidden)]
90 Ignored(&'static str),
91}
92
93#[derive(Clone, Debug, PartialEq)]
98#[non_exhaustive]
99pub enum Rewrite {
100 #[doc(hidden)]
105 OutputCapRename,
106 #[doc(hidden)]
111 DropUnboundThinking,
112 #[doc(hidden)]
117 ToolChoiceNeedsTools,
118 Stream(bool),
120 NoStream,
122 #[doc(hidden)]
127 NoBackground,
128 StreamUsage,
131 #[doc(hidden)]
136 ReasoningCiphertext(bool),
137 #[doc(hidden)]
141 CodexStore,
142 #[doc(hidden)]
146 ChatDialect(BodyRewrite),
147 #[doc(hidden)]
154 GeminiCachedContent(Option<String>),
155}
156
157pub struct BaseInput<'a> {
160 target: &'a dyn ReplayTarget,
161 request: &'a CompletionRequest,
162 cache: Option<CacheRetention>,
163 upper: &'a Map<String, Value>,
164 raw_tools: Option<Value>,
165}
166
167impl BaseInput<'_> {
168 pub fn cache(&self) -> Option<CacheRetention> {
171 self.cache
172 }
173
174 pub fn refuse_cache(&mut self, reason: impl Into<String>) -> Result<(), EncodeError> {
182 refuse(self.target, self.request, "cache", reason.into())
183 }
184
185 pub fn param(&self, key: &str) -> Option<&Value> {
190 self.upper.get(key)
191 }
192
193 pub fn raw_tools(&mut self) -> Result<Vec<Value>, EncodeError> {
201 match self.raw_tools.take() {
202 None | Some(Value::Null) => Ok(Vec::new()),
203 Some(Value::Array(tools)) => Ok(tools),
204 Some(other) => Err(EncodeError::request(format!(
205 "`additional_params.tools` must be an array, got {other}"
206 ))),
207 }
208 }
209}
210
211#[derive(Default)]
213struct Settled {
214 sends: Vec<Value>,
216 cache: Option<CacheRetention>,
218 ignored: Vec<&'static str>,
220}
221
222fn model_of<'a>(target: &'a dyn ReplayTarget, request: &'a CompletionRequest) -> &'a str {
224 request
225 .model
226 .as_deref()
227 .filter(|model| !model.is_empty())
228 .unwrap_or_else(|| target.model())
229}
230
231fn refuse(
233 target: &dyn ReplayTarget,
234 request: &CompletionRequest,
235 option: impl Into<Cow<'static, str>>,
236 reason: String,
237) -> Result<(), EncodeError> {
238 let option = option.into();
239 let provider = target.provider();
240 let model = model_of(target, request);
241 match request.options.unsupported_policy() {
242 OnUnsupported::Error => Err(EncodeError::unsupported(UnsupportedOption::new(
243 option, provider, model, reason,
244 ))),
245 OnUnsupported::Ignore => {
246 tracing::warn!(
247 option = option.as_ref(),
248 provider,
249 model,
250 reason = reason.as_str(),
251 "option skipped: the provider cannot honour it"
252 );
253 Ok(())
254 }
255 }
256}
257
258fn settle(target: &dyn ReplayTarget, request: &CompletionRequest) -> Result<Settled, EncodeError> {
261 let fields = request.options.fields();
262 let set = fields.set();
263 let map = target.map_options(request, fields);
264 let provider = target.provider();
265 let mut settled = Settled::default();
266 for ((option, mapping), set) in map.into_slots().into_iter().zip(set) {
267 let fail = |what: String| Err(EncodeError::request(format!("{provider} {what}")));
268 match (mapping, set) {
269 (Mapping::Nothing, false) => {}
270 (Mapping::Nothing, true) => {
271 return fail(format!(
272 "answered `Mapping::Nothing` for the option `{option}`, which the request sets"
273 ));
274 }
275 (_, false) => {
276 return fail(format!(
277 "answered the option `{option}`, which the request does not set; \
278 an unset option answers `Mapping::Nothing` (see `Mapping::of`)"
279 ));
280 }
281 (Mapping::Send(value), true) => {
282 if !value.is_object() {
283 return fail(format!(
284 "sent a value that is not a JSON object for the option `{option}`"
285 ));
286 }
287 settled.sends.push(value);
288 if option == "cache" {
289 settled.cache = request.options.cache;
290 }
291 }
292 (Mapping::Omit(reason), true) => {
293 tracing::debug!(
294 option,
295 provider,
296 reason,
297 "option honoured by sending nothing"
298 );
299 }
300 (Mapping::Place, true) if option == "cache" => settled.cache = request.options.cache,
301 (Mapping::Place, true) => {
302 return fail(format!(
303 "answered `Mapping::Place` for the option `{option}`; only `cache` places markers"
304 ));
305 }
306 (Mapping::Unsupported(reason), true) => {
307 refuse(target, request, option, reason)?;
308 settled.ignored.push(option);
309 }
310 }
311 }
312 Ok(settled)
313}
314
315fn clear(options: &mut GenerationOptions, ignored: &[&str]) {
317 let GenerationOptions {
318 reasoning,
319 cache,
320 service_tier,
321 verbosity,
322 parallel_tool_calls,
323 top_p,
324 seed,
325 stop,
326 on_unsupported: _,
327 } = options;
328 let off = |name: &str| ignored.contains(&name);
329 if off("reasoning") {
330 *reasoning = None;
331 }
332 if off("cache") {
333 *cache = None;
334 }
335 if off("service_tier") {
336 *service_tier = None;
337 }
338 if off("verbosity") {
339 *verbosity = None;
340 }
341 if off("parallel_tool_calls") {
342 *parallel_tool_calls = None;
343 }
344 if off("top_p") {
345 *top_p = None;
346 }
347 if off("seed") {
348 *seed = None;
349 }
350 if off("stop") {
351 stop.clear();
352 }
353}
354
355pub fn check(
368 target: &dyn ReplayTarget,
369 request: &mut CompletionRequest,
370) -> Result<(), EncodeError> {
371 unserialized(request)?;
372 let settled = settle(target, request)?;
373 clear(&mut request.options, &settled.ignored);
374 let layer = provider_layer(target, request);
375 for refusal in &layer.refused {
376 refuse(
377 target,
378 request,
379 refusal.name(target),
380 refusal.reason.clone(),
381 )?;
382 }
383 let api = target.api();
384 for refusal in layer.refused {
385 request.provider_options.remove_field(
386 target.provider(),
387 &[SHARED, api.as_str()],
388 refusal.field,
389 );
390 }
391 Ok(())
392}
393
394pub(crate) struct CatalogRefusal {
399 pub(crate) field: &'static str,
400 pub(crate) reason: String,
401}
402
403pub(crate) fn catalog_refusals(
417 target: &dyn ReplayTarget,
418 request: &CompletionRequest,
419 body: FinalBody,
420 refusals: Vec<CatalogRefusal>,
421) -> Result<FinalBody, EncodeError> {
422 let provider = target.provider();
423 let model = model_of(target, request);
424 if request.options.is_default() {
425 for refusal in &refusals {
426 tracing::debug!(
427 option = refusal.field,
428 provider,
429 model,
430 reason = refusal.reason.as_str(),
431 "catalog refusal not applied: the request sets no generation options"
432 );
433 }
434 return Ok(body);
435 }
436 let FinalBody(mut body) = body;
437 for CatalogRefusal { field, reason } in refusals {
438 if request.options.unsupported_policy() == OnUnsupported::Error {
439 return Err(EncodeError::unsupported(UnsupportedOption::new(
440 field, provider, model, reason,
441 )));
442 }
443 let raw = request
444 .additional_params
445 .as_ref()
446 .and_then(|params| params.get(field));
447 let typed = field != "tools" && raw.is_none_or(|raw| body.get(field) != Some(raw));
448 if typed {
449 body.shift_remove(field);
450 }
451 tracing::warn!(
452 option = field,
453 provider,
454 model,
455 reason = reason.as_str(),
456 "{}",
457 match typed {
458 true => "option skipped: the model's catalog entry refuses it",
459 false => "catalog refusal ignored: sent as written",
460 }
461 );
462 }
463 Ok(FinalBody(body))
464}
465
466struct Refused {
469 field: &'static str,
470 section: String,
471 reason: String,
472}
473
474impl Refused {
475 fn name(&self, target: &dyn ReplayTarget) -> String {
477 format!("{}.{}.{}", target.provider(), self.section, self.field)
478 }
479}
480
481struct ProviderLayer {
484 fields: Map<String, Value>,
485 refused: Vec<Refused>,
486}
487
488fn provider_layer(target: &dyn ReplayTarget, request: &CompletionRequest) -> ProviderLayer {
493 let provider = target.provider();
494 let api = target.api();
495 let mut layer = ProviderLayer {
496 fields: Map::new(),
497 refused: Vec::new(),
498 };
499 let Some(sections) = request.provider_options.sections(provider) else {
500 return layer;
501 };
502 let mut route = None;
503 for (name, section) in sections {
504 let Value::Object(section) = section else {
505 continue;
506 };
507 if name == SHARED {
508 deep_merge(&mut layer.fields, section.clone());
509 } else if name == api.as_str() {
510 route = Some(section);
511 } else {
512 tracing::debug!(
513 provider,
514 section = name.as_str(),
515 api = api.as_str(),
516 "provider options section skipped: the request takes another route"
517 );
518 }
519 }
520 if let Some(route) = route {
521 deep_merge(&mut layer.fields, route.clone());
522 }
523 for (field, reason) in request.provider_options.refusals(provider, target, request) {
524 layer.refuse(target, request, field, reason);
525 }
526 layer
527}
528
529impl ProviderLayer {
530 fn refuse(
534 &mut self,
535 target: &dyn ReplayTarget,
536 request: &CompletionRequest,
537 field: &'static str,
538 reason: String,
539 ) {
540 if self.fields.shift_remove(field).is_none() {
541 return;
542 }
543 let api = target.api();
544 let in_route = request
545 .provider_options
546 .sections(target.provider())
547 .and_then(|sections| sections.get(api.as_str()))
548 .and_then(Value::as_object)
549 .is_some_and(|route| route.contains_key(field));
550 let section = if in_route { api.as_str() } else { SHARED };
551 self.refused.push(Refused {
552 field,
553 section: section.to_owned(),
554 reason,
555 });
556 }
557}
558
559fn unserialized(request: &CompletionRequest) -> Result<(), EncodeError> {
564 match request.provider_options.failure() {
565 Some(error) => Err(EncodeError::request(error.clone())),
566 None => Ok(()),
567 }
568}
569
570fn raw_layer(request: &CompletionRequest) -> Result<Map<String, Value>, EncodeError> {
572 match &request.additional_params {
573 None | Some(Value::Null) => Ok(Map::new()),
574 Some(Value::Object(params)) => Ok(params.clone()),
575 Some(_) => Err(EncodeError::request(
576 "`additional_params` must be a JSON object",
577 )),
578 }
579}
580
581fn placed(raw: Map<String, Value>, raw_at: RawAt) -> Map<String, Value> {
583 match raw_at {
584 RawAt::Top => raw,
585 RawAt::Under(_) if raw.is_empty() => Map::new(),
586 RawAt::Under(pointer) => {
587 let mut value = Value::Object(raw);
588 for key in pointer.rsplit('/').filter(|key| !key.is_empty()) {
589 value = Value::Object(Map::from_iter([(key.to_owned(), value)]));
590 }
591 match value {
592 Value::Object(map) => map,
593 _ => Map::new(),
594 }
595 }
596 RawAt::Split { top, rest } => {
597 let mut body = Map::new();
598 let mut under = Map::new();
599 for (key, value) in raw {
600 if key == rest || top.contains(&key.as_str()) {
601 deep_merge(&mut body, Map::from_iter([(key, value)]));
602 } else {
603 under.insert(key, value);
604 }
605 }
606 if !under.is_empty() {
607 deep_merge(
608 &mut body,
609 Map::from_iter([(rest.to_owned(), Value::Object(under))]),
610 );
611 }
612 body
613 }
614 RawAt::Ignored(warning) => {
615 if !raw.is_empty() {
616 tracing::warn!("{warning}");
617 }
618 Map::new()
619 }
620 }
621}
622
623pub(crate) fn deep_merge(body: &mut Map<String, Value>, upper: Map<String, Value>) {
626 for (key, value) in upper {
627 match (body.get_mut(&key), value) {
628 (Some(Value::Object(lower)), Value::Object(upper)) => deep_merge(lower, upper),
629 (_, value) => {
630 body.insert(key, value);
631 }
632 }
633 }
634}
635
636fn upper_layers(
639 sends: &[Value],
640 provider: &Map<String, Value>,
641 raw: &Map<String, Value>,
642) -> Map<String, Value> {
643 let mut upper = Map::new();
644 for send in sends {
645 if let Value::Object(send) = send {
646 deep_merge(&mut upper, send.clone());
647 }
648 }
649 deep_merge(&mut upper, provider.clone());
650 deep_merge(&mut upper, raw.clone());
651 upper
652}
653
654pub fn param(target: &dyn ReplayTarget, request: &CompletionRequest, key: &str) -> Option<Value> {
662 let map = target.map_options(request, request.options.fields());
663 let mut upper = Map::new();
664 for (_, mapping) in map.into_slots() {
665 if let Mapping::Send(Value::Object(send)) = mapping {
666 deep_merge(&mut upper, send);
667 }
668 }
669 deep_merge(&mut upper, provider_layer(target, request).fields);
670 if let Some(Value::Object(raw)) = &request.additional_params {
671 deep_merge(&mut upper, raw.clone());
672 }
673 upper.shift_remove(key)
674}
675
676pub fn request_params(
692 target: &dyn ReplayTarget,
693 request: &CompletionRequest,
694 base: impl FnOnce(&mut BaseInput<'_>) -> Result<Map<String, Value>, EncodeError>,
695 raw_at: RawAt,
696 rewrites: &[Rewrite],
697) -> Result<FinalBody, EncodeError> {
698 unserialized(request)?;
699 let settled = settle(target, request)?;
700 let mut provider = provider_layer(target, request);
701 if rewrites.contains(&Rewrite::NoBackground) {
704 provider.refuse(
705 target,
706 request,
707 "background",
708 "a WebSocket session takes no `background`".to_owned(),
709 );
710 }
711 for refusal in provider.refused {
712 refuse(target, request, refusal.name(target), refusal.reason)?;
713 }
714 let provider = provider.fields;
715 let mut raw = raw_layer(request)?;
716 let raw_tools = match raw_at {
717 RawAt::Top | RawAt::Split { .. } => raw.shift_remove("tools"),
718 RawAt::Under(_) | RawAt::Ignored(_) => None,
719 };
720 let mut handles = Vec::new();
721 if rewrites
722 .iter()
723 .any(|rewrite| matches!(rewrite, Rewrite::GeminiCachedContent(_)))
724 {
725 for spelling in ["cachedContent", "cached_content"] {
726 match raw.shift_remove(spelling) {
727 None => {}
728 Some(Value::String(name)) => handles.push(name),
729 Some(other) => {
730 return Err(EncodeError::request(format!(
731 "Gemini `additional_params.{spelling}` should be a string, got {other}"
732 )));
733 }
734 }
735 }
736 if raw.get("generationConfig").is_some_and(Value::is_null) {
739 raw.shift_remove("generationConfig");
740 }
741 }
742 let raw = placed(raw, raw_at);
743 let upper = upper_layers(&settled.sends, &provider, &raw);
744 let mut input = BaseInput {
745 target,
746 request,
747 cache: settled.cache,
748 upper: &upper,
749 raw_tools,
750 };
751 let mut body = base(&mut input)?;
752 for send in settled.sends {
753 if let Value::Object(send) = send {
754 deep_merge(&mut body, send);
755 }
756 }
757 deep_merge(&mut body, provider);
758 deep_merge(&mut body, raw);
759 for rewrite in rewrites {
760 apply(rewrite, &mut body, &handles)?;
761 }
762 Ok(FinalBody(body))
763}
764
765fn apply(
767 rewrite: &Rewrite,
768 body: &mut Map<String, Value>,
769 handles: &[String],
770) -> Result<(), EncodeError> {
771 match rewrite {
772 Rewrite::OutputCapRename => {
773 let reasoning = body
776 .get("model")
777 .and_then(Value::as_str)
778 .is_some_and(|model| {
779 !model.contains('/')
780 && crate::providers::openai::options::reasons(model) == Some(true)
781 });
782 if reasoning && let Some(max_tokens) = body.shift_remove("max_tokens") {
783 body.entry("max_completion_tokens").or_insert(max_tokens);
784 }
785 }
786 Rewrite::DropUnboundThinking => {
787 let adaptive = body.get("thinking").is_none_or(|thinking| {
788 thinking.get("type").and_then(Value::as_str) == Some("adaptive")
789 });
790 if adaptive {
791 crate::providers::anthropic::completion::drop_unbound_thinking(body);
792 }
793 }
794 Rewrite::ToolChoiceNeedsTools => {
795 let has_tools = body
796 .get("tools")
797 .and_then(Value::as_array)
798 .is_some_and(|tools| !tools.is_empty());
799 if has_tools {
800 body.entry("tool_choice")
801 .or_insert_with(|| serde_json::json!({ "type": "auto" }));
802 } else {
803 body.shift_remove("tool_choice");
804 }
805 }
806 Rewrite::Stream(stream) => {
807 body.insert("stream".to_owned(), Value::Bool(*stream));
808 }
809 Rewrite::NoStream => {
810 body.shift_remove("stream");
811 }
812 Rewrite::NoBackground => {
813 body.shift_remove("background");
814 }
815 Rewrite::StreamUsage => {
816 if let Some(options) = body
817 .entry("stream_options")
818 .or_insert_with(|| Value::Object(Map::new()))
819 .as_object_mut()
820 {
821 options.entry("include_usage").or_insert(Value::Bool(true));
822 }
823 }
824 Rewrite::ReasoningCiphertext(always) => {
825 let wanted = *always
826 || body.get("reasoning").is_some()
827 || body.get("store") == Some(&Value::Bool(false));
828 if wanted {
829 crate::providers::openai::responses_api::include_ciphertext(body);
830 }
831 }
832 Rewrite::CodexStore => {
833 body.insert("store".to_owned(), Value::Bool(false));
834 }
835 Rewrite::ChatDialect(kind) => {
836 crate::providers::openai::wire::chat::rewrite_body(*kind, body)?;
837 }
838 Rewrite::GeminiCachedContent(handle) => {
839 for name in handles.iter().chain(handle) {
840 crate::providers::gemini::completion::with_cached_content(body, name)?;
841 }
842 }
843 }
844 Ok(())
845}