1use crate::limits::ExtensionLimits;
4use serde::{Deserialize, Serialize};
5use std::collections::{BTreeMap, BTreeSet};
6use std::time::Duration;
7use thiserror::Error;
8
9#[derive(
11 Clone, Copy, Debug, Default, PartialEq, Eq, PartialOrd, Ord, Hash, Serialize, Deserialize,
12)]
13pub enum ContinuationPolicy {
14 InlineToolContinuation,
16 #[default]
18 CallerControlled,
19}
20
21#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash, Serialize, Deserialize)]
23pub enum ReasoningEffort {
24 Low,
26 Medium,
28 High,
30}
31
32#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
34pub enum ResponseFormat {
35 Text,
37 JsonObject,
39}
40
41#[derive(Clone, Debug, PartialEq, Eq, PartialOrd, Ord, Hash, Serialize, Deserialize)]
43pub struct ExtensionKey(String);
44
45impl ExtensionKey {
46 pub fn try_new(value: impl Into<String>, max_bytes: usize) -> Result<Self, ConfigError> {
48 let s = value.into();
49 if s.is_empty() {
50 return Err(ConfigError::EmptyExtensionKey);
51 }
52 if s.len() > max_bytes {
53 return Err(ConfigError::ExtensionKeyTooLong {
54 bytes: s.len(),
55 max: max_bytes,
56 });
57 }
58 if s.chars().any(|c| c.is_control()) {
59 return Err(ConfigError::ControlCharacter);
60 }
61 if !s.contains('.') {
62 return Err(ConfigError::ExtensionKeyMissingNamespace);
63 }
64 Ok(Self(s))
65 }
66
67 pub fn as_str(&self) -> &str {
69 &self.0
70 }
71}
72
73#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
75pub struct VersionedExtension {
76 pub version: u16,
78 pub value: serde_json::Value,
80}
81
82#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
84pub struct InvocationConfig {
85 pub model: Option<String>,
87 pub temperature: Option<f32>,
89 pub reasoning_effort: Option<ReasoningEffort>,
91 pub max_output_tokens: Option<u32>,
93 pub stop: Vec<String>,
95 pub response_format: Option<ResponseFormat>,
97 pub continuation_policy: ContinuationPolicy,
99 pub deadline: Option<Duration>,
101 pub extensions: BTreeMap<ExtensionKey, VersionedExtension>,
103}
104
105impl Default for InvocationConfig {
106 fn default() -> Self {
107 Self {
108 model: None,
109 temperature: None,
110 reasoning_effort: None,
111 max_output_tokens: None,
112 stop: Vec::new(),
113 response_format: None,
114 continuation_policy: ContinuationPolicy::CallerControlled,
115 deadline: None,
116 extensions: BTreeMap::new(),
117 }
118 }
119}
120
121#[derive(Clone, Debug, PartialEq, Serialize, Deserialize, Default)]
123pub struct SessionConfig {
124 pub specialist_profile: Option<String>,
126 pub mode: Option<String>,
128 pub permission_profile: Option<String>,
130 pub extensions: BTreeMap<ExtensionKey, VersionedExtension>,
132 pub connector_ref: Option<String>,
142}
143
144#[derive(Clone, Debug, PartialEq, Serialize, Deserialize, Default)]
146pub struct ChannelDefaults {
147 pub model: Option<String>,
149 pub temperature: Option<f32>,
151 pub reasoning_effort: Option<ReasoningEffort>,
153 pub max_output_tokens: Option<u32>,
155 pub stop: Vec<String>,
157 pub response_format: Option<ResponseFormat>,
159 pub continuation_policy: ContinuationPolicy,
161 pub extensions: BTreeMap<ExtensionKey, VersionedExtension>,
163}
164
165#[derive(Clone, Debug, PartialEq, Eq, Default, Serialize, Deserialize)]
167pub struct OptionPolicy {
168 pub supported_invocation: BTreeSet<ConfigOption>,
170 pub session_immutable: BTreeSet<ConfigOption>,
172 pub allowed_extension_keys: BTreeSet<ExtensionKey>,
175}
176
177impl OptionPolicy {
178 pub fn direct_llm() -> Self {
180 let mut p = Self::default();
181 p.supported_invocation.extend([
182 ConfigOption::Model,
183 ConfigOption::Temperature,
184 ConfigOption::ReasoningEffort,
185 ConfigOption::MaxOutputTokens,
186 ConfigOption::Stop,
187 ConfigOption::ResponseFormat,
188 ConfigOption::ContinuationPolicy,
189 ConfigOption::Deadline,
190 ConfigOption::Extensions,
191 ]);
192 p
193 }
194
195 pub fn external_agent() -> Self {
197 let mut p = Self::default();
198 p.supported_invocation.extend([
199 ConfigOption::ContinuationPolicy,
200 ConfigOption::Deadline,
201 ConfigOption::Extensions,
202 ]);
203 p
204 }
205
206 pub fn with_extension_keys(mut self, keys: impl IntoIterator<Item = ExtensionKey>) -> Self {
208 self.allowed_extension_keys.extend(keys);
209 self
210 }
211}
212
213#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash, Serialize, Deserialize)]
215pub enum ConfigOption {
216 Model,
218 Temperature,
220 ReasoningEffort,
222 MaxOutputTokens,
224 Stop,
226 ResponseFormat,
228 ContinuationPolicy,
230 Deadline,
232 Extensions,
234}
235
236#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
238pub struct EffectiveConfig {
239 pub model: Option<String>,
241 pub temperature: Option<f32>,
243 pub reasoning_effort: Option<ReasoningEffort>,
245 pub max_output_tokens: Option<u32>,
247 pub stop: Vec<String>,
249 pub response_format: Option<ResponseFormat>,
251 pub continuation_policy: ContinuationPolicy,
253 pub deadline: Option<Duration>,
255 pub extensions: BTreeMap<ExtensionKey, VersionedExtension>,
257 pub session: SessionConfig,
259}
260
261pub fn merge_effective_config(
263 defaults: &ChannelDefaults,
264 session: Option<&SessionConfig>,
265 attached_session: Option<&SessionConfig>,
266 invocation: &InvocationConfig,
267 policy: &OptionPolicy,
268 extension_limits: &ExtensionLimits,
269) -> Result<EffectiveConfig, ConfigError> {
270 validate_extensions(&defaults.extensions, extension_limits, policy)?;
271 if let Some(s) = session {
272 validate_session_labels(s)?;
273 validate_extensions(&s.extensions, extension_limits, policy)?;
274 }
275 if let Some(s) = attached_session {
276 validate_session_labels(s)?;
277 validate_extensions(&s.extensions, extension_limits, policy)?;
278 }
279 validate_extensions(&invocation.extensions, extension_limits, policy)?;
280 validate_invocation_strings(invocation)?;
281
282 if let (Some(requested), Some(attached)) = (session, attached_session) {
284 check_session_match(requested, attached, policy)?;
285 }
286
287 let mut effective = EffectiveConfig {
288 model: defaults.model.clone(),
289 temperature: defaults.temperature,
290 reasoning_effort: defaults.reasoning_effort,
291 max_output_tokens: defaults.max_output_tokens,
292 stop: defaults.stop.clone(),
293 response_format: defaults.response_format.clone(),
294 continuation_policy: defaults.continuation_policy,
295 deadline: None,
296 extensions: defaults.extensions.clone(),
297 session: session.cloned().unwrap_or_default(),
298 };
299
300 if let Some(s) = session {
302 for (k, v) in &s.extensions {
303 effective.extensions.insert(k.clone(), v.clone());
304 }
305 }
306
307 apply_invocation(&mut effective, invocation, policy)?;
308
309 validate_extensions(&effective.extensions, extension_limits, policy)?;
311 Ok(effective)
312}
313
314fn apply_invocation(
315 effective: &mut EffectiveConfig,
316 invocation: &InvocationConfig,
317 policy: &OptionPolicy,
318) -> Result<(), ConfigError> {
319 if invocation.model.is_some() {
320 require_supported(policy, ConfigOption::Model)?;
321 effective.model = invocation.model.clone();
322 }
323 if invocation.temperature.is_some() {
324 require_supported(policy, ConfigOption::Temperature)?;
325 if let Some(t) = invocation.temperature {
326 if !(0.0..=2.0).contains(&t) {
327 return Err(ConfigError::InvalidNumeric("temperature"));
328 }
329 }
330 effective.temperature = invocation.temperature;
331 }
332 if invocation.reasoning_effort.is_some() {
333 require_supported(policy, ConfigOption::ReasoningEffort)?;
334 effective.reasoning_effort = invocation.reasoning_effort;
335 }
336 if invocation.max_output_tokens.is_some() {
337 require_supported(policy, ConfigOption::MaxOutputTokens)?;
338 effective.max_output_tokens = invocation.max_output_tokens;
339 }
340 if !invocation.stop.is_empty() {
341 require_supported(policy, ConfigOption::Stop)?;
342 effective.stop = invocation.stop.clone();
343 }
344 if invocation.response_format.is_some() {
345 require_supported(policy, ConfigOption::ResponseFormat)?;
346 effective.response_format = invocation.response_format.clone();
347 }
348 require_supported(policy, ConfigOption::ContinuationPolicy)?;
350 effective.continuation_policy = invocation.continuation_policy;
351
352 if invocation.deadline.is_some() {
353 require_supported(policy, ConfigOption::Deadline)?;
354 effective.deadline = invocation.deadline;
355 }
356 if !invocation.extensions.is_empty() {
357 require_supported(policy, ConfigOption::Extensions)?;
358 for (k, v) in &invocation.extensions {
359 effective.extensions.insert(k.clone(), v.clone());
360 }
361 }
362 Ok(())
363}
364
365fn require_supported(policy: &OptionPolicy, opt: ConfigOption) -> Result<(), ConfigError> {
366 if policy.supported_invocation.contains(&opt) {
367 Ok(())
368 } else {
369 Err(ConfigError::UnsupportedOption(opt))
370 }
371}
372
373fn check_session_match(
374 requested: &SessionConfig,
375 attached: &SessionConfig,
376 policy: &OptionPolicy,
377) -> Result<(), ConfigError> {
378 if policy.session_immutable.contains(&ConfigOption::Model) {
379 }
381 if requested.specialist_profile != attached.specialist_profile
383 && (requested.specialist_profile.is_some() || attached.specialist_profile.is_some())
384 {
385 return Err(ConfigError::ImmutableSessionMismatch("specialist_profile"));
386 }
387 if requested.mode != attached.mode && (requested.mode.is_some() || attached.mode.is_some()) {
388 return Err(ConfigError::ImmutableSessionMismatch("mode"));
389 }
390 if requested.permission_profile != attached.permission_profile
391 && (requested.permission_profile.is_some() || attached.permission_profile.is_some())
392 {
393 return Err(ConfigError::ImmutableSessionMismatch("permission_profile"));
394 }
395 for (k, v) in &requested.extensions {
396 if let Some(existing) = attached.extensions.get(k) {
397 if existing != v {
398 return Err(ConfigError::ImmutableSessionMismatch("extension"));
399 }
400 }
401 }
402 Ok(())
403}
404
405fn validate_session_labels(session: &SessionConfig) -> Result<(), ConfigError> {
406 for label in [
407 &session.specialist_profile,
408 &session.mode,
409 &session.permission_profile,
410 ]
411 .into_iter()
412 .flatten()
413 {
414 if label.is_empty() || label.len() > 128 || label.chars().any(|c| c.is_control()) {
415 return Err(ConfigError::InvalidSessionLabel);
416 }
417 }
418 Ok(())
419}
420
421fn validate_invocation_strings(invocation: &InvocationConfig) -> Result<(), ConfigError> {
422 if let Some(m) = &invocation.model {
423 if m.is_empty() || m.len() > 256 || m.chars().any(|c| c.is_control()) {
424 return Err(ConfigError::InvalidModel);
425 }
426 }
427 for s in &invocation.stop {
428 if s.is_empty() || s.len() > 64 || s.chars().any(|c| c.is_control()) {
429 return Err(ConfigError::InvalidStop);
430 }
431 }
432 Ok(())
433}
434
435fn validate_extensions(
436 map: &BTreeMap<ExtensionKey, VersionedExtension>,
437 limits: &ExtensionLimits,
438 policy: &OptionPolicy,
439) -> Result<(), ConfigError> {
440 if map.len() > limits.max_keys {
441 return Err(ConfigError::TooManyExtensions {
442 count: map.len(),
443 max: limits.max_keys,
444 });
445 }
446 if !map.is_empty() && policy.allowed_extension_keys.is_empty() {
448 let first = map
449 .keys()
450 .next()
451 .map(|k| k.as_str().to_string())
452 .unwrap_or_default();
453 return Err(ConfigError::UnknownExtension(first));
454 }
455 let mut total = 0usize;
456 for (k, v) in map {
457 if k.as_str().len() > limits.max_key_bytes {
458 return Err(ConfigError::ExtensionKeyTooLong {
459 bytes: k.as_str().len(),
460 max: limits.max_key_bytes,
461 });
462 }
463 if !policy.allowed_extension_keys.contains(k) {
464 return Err(ConfigError::UnknownExtension(k.as_str().to_string()));
465 }
466 let depth = json_depth(&v.value);
467 if depth > limits.max_value_depth {
468 return Err(ConfigError::ExtensionTooDeep {
469 depth,
470 max: limits.max_value_depth,
471 });
472 }
473 let encoded =
474 serde_json::to_vec(&v.value).map_err(|_| ConfigError::ExtensionEncodeFailed)?;
475 total = total
476 .saturating_add(encoded.len())
477 .saturating_add(k.as_str().len());
478 }
479 if total > limits.max_serialized_bytes {
480 return Err(ConfigError::ExtensionsTooLarge {
481 bytes: total,
482 max: limits.max_serialized_bytes,
483 });
484 }
485 Ok(())
486}
487
488fn json_depth(value: &serde_json::Value) -> u32 {
489 match value {
490 serde_json::Value::Array(items) => 1 + items.iter().map(json_depth).max().unwrap_or(0),
491 serde_json::Value::Object(map) => 1 + map.values().map(json_depth).max().unwrap_or(0),
492 _ => 1,
493 }
494}
495
496#[derive(Clone, Debug, Error, PartialEq)]
498pub enum ConfigError {
499 #[error("extension key must be non-empty")]
501 EmptyExtensionKey,
502 #[error("extension key must be namespaced (contain '.')")]
504 ExtensionKeyMissingNamespace,
505 #[error("extension key bytes {bytes} exceeds max {max}")]
507 ExtensionKeyTooLong {
508 bytes: usize,
510 max: usize,
512 },
513 #[error("configuration string must not contain control characters")]
515 ControlCharacter,
516 #[error("extension count {count} exceeds max {max}")]
518 TooManyExtensions {
519 count: usize,
521 max: usize,
523 },
524 #[error("extension depth {depth} exceeds max {max}")]
526 ExtensionTooDeep {
527 depth: u32,
529 max: u32,
531 },
532 #[error("extensions serialized bytes {bytes} exceed max {max}")]
534 ExtensionsTooLarge {
535 bytes: usize,
537 max: usize,
539 },
540 #[error("unknown extension key {0}")]
542 UnknownExtension(String),
543 #[error("extension JSON encode failed")]
545 ExtensionEncodeFailed,
546 #[error("unsupported configuration option: {0:?}")]
548 UnsupportedOption(ConfigOption),
549 #[error("invalid numeric value for {0}")]
551 InvalidNumeric(&'static str),
552 #[error("invalid model string")]
554 InvalidModel,
555 #[error("invalid stop sequence")]
557 InvalidStop,
558 #[error("invalid session configuration label")]
560 InvalidSessionLabel,
561 #[error("immutable session setting mismatch: {0}")]
563 ImmutableSessionMismatch(&'static str),
564}
565
566#[cfg(test)]
567mod tests {
568 use super::*;
569
570 fn open_policy() -> OptionPolicy {
571 let mut p = OptionPolicy::default();
572 p.supported_invocation.extend([
573 ConfigOption::Model,
574 ConfigOption::Temperature,
575 ConfigOption::ContinuationPolicy,
576 ConfigOption::Deadline,
577 ConfigOption::Extensions,
578 ]);
579 p
580 }
581
582 #[test]
583 fn merge_precedence_invocation_over_defaults() {
584 let defaults = ChannelDefaults {
585 model: Some("base".into()),
586 temperature: Some(0.2),
587 continuation_policy: ContinuationPolicy::CallerControlled,
588 ..Default::default()
589 };
590 let inv = InvocationConfig {
591 model: Some("override".into()),
592 temperature: Some(0.7),
593 continuation_policy: ContinuationPolicy::InlineToolContinuation,
594 ..Default::default()
595 };
596 let eff = merge_effective_config(
597 &defaults,
598 None,
599 None,
600 &inv,
601 &open_policy(),
602 &ExtensionLimits::default(),
603 )
604 .unwrap();
605 assert_eq!(eff.model.as_deref(), Some("override"));
606 assert_eq!(eff.temperature, Some(0.7));
607 assert_eq!(
608 eff.continuation_policy,
609 ContinuationPolicy::InlineToolContinuation
610 );
611 }
612
613 #[test]
614 fn immutable_session_mismatch_fails() {
615 let requested = SessionConfig {
616 mode: Some("agent".into()),
617 ..Default::default()
618 };
619 let attached = SessionConfig {
620 mode: Some("ask".into()),
621 ..Default::default()
622 };
623 let err = merge_effective_config(
624 &ChannelDefaults::default(),
625 Some(&requested),
626 Some(&attached),
627 &InvocationConfig {
628 continuation_policy: ContinuationPolicy::CallerControlled,
629 ..Default::default()
630 },
631 &open_policy(),
632 &ExtensionLimits::default(),
633 )
634 .unwrap_err();
635 assert!(matches!(err, ConfigError::ImmutableSessionMismatch("mode")));
636 }
637
638 #[test]
639 fn extension_bounds() {
640 let limits = ExtensionLimits {
641 max_keys: 1,
642 max_key_bytes: 32,
643 max_value_depth: 2,
644 max_serialized_bytes: 64,
645 };
646 let k = ExtensionKey::try_new("ns.a", limits.max_key_bytes).unwrap();
647 let k2 = ExtensionKey::try_new("ns.b", limits.max_key_bytes).unwrap();
648 let mut inv = InvocationConfig::default();
649 inv.extensions.insert(
650 k,
651 VersionedExtension {
652 version: 1,
653 value: serde_json::json!(1),
654 },
655 );
656 inv.extensions.insert(
657 k2,
658 VersionedExtension {
659 version: 1,
660 value: serde_json::json!(2),
661 },
662 );
663 let policy = open_policy();
664 let err = merge_effective_config(
666 &ChannelDefaults::default(),
667 None,
668 None,
669 &inv,
670 &policy,
671 &limits,
672 )
673 .unwrap_err();
674 assert!(matches!(err, ConfigError::TooManyExtensions { .. }));
675 }
676
677 #[test]
679 fn empty_extension_allowlist_denies() {
680 let limits = ExtensionLimits {
681 max_keys: 8,
682 max_key_bytes: 32,
683 max_value_depth: 2,
684 max_serialized_bytes: 256,
685 };
686 let k = ExtensionKey::try_new("ns.secret", limits.max_key_bytes).unwrap();
687 let mut inv = InvocationConfig::default();
688 inv.extensions.insert(
689 k,
690 VersionedExtension {
691 version: 1,
692 value: serde_json::json!({"x": 1}),
693 },
694 );
695 let mut policy = open_policy();
696 policy.allowed_extension_keys.clear();
697 let err = merge_effective_config(
698 &ChannelDefaults::default(),
699 None,
700 None,
701 &inv,
702 &policy,
703 &limits,
704 )
705 .unwrap_err();
706 assert!(matches!(err, ConfigError::UnknownExtension(_)));
707 }
708}