1use std::collections::{HashMap, HashSet, hash_map::DefaultHasher};
21use std::hash::{Hash, Hasher};
22use std::sync::atomic::{AtomicBool, Ordering};
23
24use async_trait::async_trait;
25use parking_lot::Mutex;
26use serde::Deserialize;
27use switchyard_protocol::{ContentBlock, Message, ModelId, Request, Role};
28
29use crate::core::algorithm::{Driver, RoutingIdentity};
30use crate::core::classifier::{Classification, Classifier, Score};
31use crate::core::processor::{Event, Processor};
32
33const MAX_ASSIGNMENTS: usize = 4096;
36
37#[derive(Clone, Copy, Debug, Default, Deserialize, PartialEq, Eq)]
39#[serde(rename_all = "snake_case")]
40pub enum ClassifyTrigger {
41 #[default]
43 EveryRequest,
44 UserTurn,
46 NewSession,
48}
49
50fn is_user_turn(message: &Message) -> bool {
53 message.role == Role::User
54 && !message
55 .content
56 .iter()
57 .all(|block| matches!(block, ContentBlock::ToolResult(_)))
58}
59
60fn has_new_user_turn(messages: &[Message]) -> bool {
61 messages.last().is_some_and(is_user_turn)
62}
63
64#[derive(Default)]
75pub struct AffinityRouter {
76 latch_only: Option<HashSet<ModelId>>,
78 subagents_only: bool,
80 message_hash_fallback: bool,
82 release_on_user_turn: bool,
84 assignments: Mutex<HashMap<RoutingIdentity, ModelId>>,
89 unkeyed_warning_emitted: AtomicBool,
91}
92
93impl AffinityRouter {
94 pub fn new() -> Self {
96 Self::default()
97 }
98
99 pub fn for_subagents() -> Self {
104 Self {
105 subagents_only: true,
106 ..Self::default()
107 }
108 }
109
110 pub fn with_message_hash_fallback(mut self) -> Self {
112 self.message_hash_fallback = true;
113 self
114 }
115
116 pub fn with_release_on_user_turn(mut self) -> Self {
118 self.release_on_user_turn = true;
119 self
120 }
121
122 pub fn with_latch_only(mut self, models: impl IntoIterator<Item = impl Into<ModelId>>) -> Self {
125 self.latch_only = Some(models.into_iter().map(Into::into).collect());
126 self
127 }
128
129 fn should_latch(&self, model: &str) -> bool {
131 self.latch_only
132 .as_ref()
133 .is_none_or(|set| set.contains(model))
134 }
135
136 fn affinity_key(&self, request: &Request) -> Option<RoutingIdentity> {
138 let metadata = request.metadata.as_ref();
139 let is_subagent = metadata.is_some_and(|metadata| metadata.is_subagent);
140 let is_subagent_work = metadata.is_some_and(|metadata| metadata.is_subagent_work());
141 if self.subagents_only && !is_subagent_work {
143 return None;
144 }
145
146 let key = match (RoutingIdentity::from_request(request), is_subagent) {
147 (Some(identity), _) => Some(identity),
148 (None, true) => None,
151 (None, false) if self.message_hash_fallback => {
152 first_user_message_hash(request).map(|hash| {
153 tracing::debug!(affinity_key = %hash, "affinity using message hash fallback");
154 RoutingIdentity::Session(hash)
155 })
156 }
157 (None, false) => None,
158 };
159 if key.is_none() && self.should_warn_unkeyed() {
162 tracing::warn!(
163 target: "libsy",
164 is_subagent,
165 message_hash_fallback = self.message_hash_fallback,
166 "affinity is enabled but this request carries no usable identity, so no \
167 affinity is applied; root requests need a session id or message-hash \
168 fallback with usable first-user text, and child requests need both session \
169 and agent ids"
170 );
171 }
172 key
173 }
174
175 fn should_warn_unkeyed(&self) -> bool {
177 !self.unkeyed_warning_emitted.swap(true, Ordering::Relaxed)
178 }
179}
180
181#[async_trait]
182impl<S> Processor<S> for AffinityRouter
183where
184 S: Send + 'static,
185{
186 async fn process(&self, _state: &mut S, event: Event<'_>) -> crate::Result<()> {
187 if let Event::Decision {
188 request,
189 selected_model_id,
190 } = event
191 && let Some(key) = self.affinity_key(request)
192 {
193 let mut assignments = self.assignments.lock();
194 let writable = self.release_on_user_turn || !assignments.contains_key(&key);
195 if self.should_latch(selected_model_id) && writable {
196 evict_if_full(&mut assignments);
197 assignments.insert(key, selected_model_id.clone());
198 }
199 }
200 Ok(())
201 }
202}
203
204fn first_user_message_hash(request: &Request) -> Option<String> {
208 let message = request
209 .llm_request
210 .messages
211 .iter()
212 .find(|message| message.role == Role::User)?;
213 let mut hasher = DefaultHasher::new();
214 message.text_content("")?.hash(&mut hasher);
215 Some(format!("{:016x}", hasher.finish()))
216}
217
218#[async_trait]
219impl<S> Classifier<S> for AffinityRouter
220where
221 S: Send + 'static,
222{
223 async fn score(
224 &self,
225 _state: &mut S,
226 request: &mut Request,
227 _driver: Option<&Driver>,
228 ) -> crate::Result<(Classification, Option<switchyard_protocol::Response>)> {
229 let Some(key) = self.affinity_key(request) else {
230 return Ok((Classification::Scores(Vec::new()), None));
231 };
232 if self.release_on_user_turn && has_new_user_turn(&request.llm_request.messages) {
235 return Ok((Classification::Scores(Vec::new()), None));
236 }
237 let assigned = self.assignments.lock().get(&key).cloned();
238 Ok((
239 Classification::Scores(match assigned {
240 Some(target) => vec![Score {
241 confidence: 1.0,
242 target,
243 }],
244 None => Vec::new(),
245 }),
246 None,
247 ))
248 }
249}
250
251fn evict_if_full(assignments: &mut HashMap<RoutingIdentity, ModelId>) {
253 if assignments.len() >= MAX_ASSIGNMENTS
254 && let Some(evicted) = assignments.keys().next().cloned()
255 {
256 assignments.remove(&evicted);
257 }
258}
259
260#[cfg(test)]
261mod tests {
262 use super::*;
263
264 use std::sync::Arc;
265
266 use switchyard_protocol::{LlmRequest, Metadata, ToolResult, text_request};
267
268 type BoxErr = Box<dyn std::error::Error + Send + Sync>;
270
271 fn fixed_model(target: &str) -> ModelId {
272 ModelId::from(target)
273 }
274
275 fn request(metadata: Metadata) -> Request {
276 Request {
277 llm_request: text_request(Some("auto".to_string()), "hi"),
278 raw_request: None,
279 metadata: Some(metadata),
280 }
281 }
282
283 fn task_request(
284 metadata: Option<Metadata>,
285 first_user: &str,
286 follow_up: Option<&str>,
287 ) -> Request {
288 let mut messages = vec![
289 Message::text(Role::System, "follow repository instructions"),
290 Message::text(Role::User, first_user),
291 ];
292 if let Some(follow_up) = follow_up {
293 messages.push(Message::text(Role::Assistant, "I will inspect the code."));
294 messages.push(Message::text(Role::User, follow_up));
295 }
296 Request {
297 llm_request: LlmRequest {
298 model: Some("auto".to_string()),
299 messages,
300 ..LlmRequest::default()
301 },
302 raw_request: None,
303 metadata,
304 }
305 }
306
307 fn session(session_id: &str, agent_id: &str) -> Metadata {
308 Metadata {
309 session_id: Some(session_id.to_string()),
310 agent_id: Some(agent_id.to_string()),
311 ..Metadata::default()
312 }
313 }
314
315 fn subagent(agent_id: &str, task_id: &str) -> Metadata {
316 Metadata {
317 session_id: Some("session-1".to_string()),
318 agent_id: Some(agent_id.to_string()),
319 task_id: Some(task_id.to_string()),
320 is_subagent: true,
321 is_delegated_work: true,
322 ..Metadata::default()
323 }
324 }
325
326 async fn retain(
328 router: &AffinityRouter,
329 state: &mut (),
330 request: &mut Request,
331 model: &'static str,
332 ) -> Result<(), BoxErr> {
333 let selected_model_id = ModelId::from(model);
334 router
335 .process(
336 state,
337 Event::Decision {
338 request,
339 selected_model_id: &selected_model_id,
340 },
341 )
342 .await?;
343 Ok(())
344 }
345
346 async fn scores(
348 classifier: &dyn Classifier,
349 state: &mut (),
350 request: &mut Request,
351 ) -> Result<Vec<Score>, BoxErr> {
352 match classifier.score(state, request, None).await?.0 {
353 Classification::Scores(scores) => Ok(scores),
354 Classification::Ambiguous(_) => Err("affinity never returns ambiguous scores".into()),
355 }
356 }
357
358 #[tokio::test]
359 async fn session_retains_first_model_across_requests() -> Result<(), BoxErr> {
360 let router = AffinityRouter::new();
361 let mut state = ();
362
363 let mut first = request(session("session-1", "agent-a"));
364 retain(&router, &mut state, &mut first, "model-a").await?;
365
366 let mut second = request(session("session-1", "agent-b"));
368 let scores = scores(&router, &mut state, &mut second).await?;
369 assert_eq!(scores.len(), 1);
370 assert_eq!(scores[0].confidence, 1.0);
371 assert_eq!(scores[0].target, "model-a");
372 Ok(())
373 }
374
375 #[tokio::test]
376 async fn subagent_only_retains_children_without_latching_root_traffic() -> Result<(), BoxErr> {
377 let router = AffinityRouter::for_subagents();
378 let mut state = ();
379
380 let mut root = request(session("session-1", "root-agent"));
381 retain(&router, &mut state, &mut root, "model-a").await?;
382 assert!(scores(&router, &mut state, &mut root).await?.is_empty());
383
384 let mut first_child_turn = request(subagent("child-1", "task-1"));
385 retain(&router, &mut state, &mut first_child_turn, "model-b").await?;
386 let mut later_child_turn = request(subagent("child-1", "task-2"));
387 let scores = scores(&router, &mut state, &mut later_child_turn).await?;
388 assert_eq!(
389 scores.first().map(|score| score.target.as_str()),
390 Some("model-b")
391 );
392 Ok(())
393 }
394
395 #[tokio::test]
396 async fn first_decision_wins() -> Result<(), BoxErr> {
397 let router = AffinityRouter::new();
398 let mut state = ();
399
400 let mut req = request(session("session-1", "agent-a"));
401 retain(&router, &mut state, &mut req, "model-a").await?;
402 retain(&router, &mut state, &mut req, "model-b").await?;
404
405 let scores = scores(&router, &mut state, &mut req).await?;
406 assert_eq!(
407 scores.first().map(|score| score.target.as_str()),
408 Some("model-a")
409 );
410 Ok(())
411 }
412
413 #[tokio::test]
414 async fn subagent_is_keyed_by_agent_not_task() -> Result<(), BoxErr> {
415 let router = AffinityRouter::new();
416 let mut state = ();
417
418 let mut first = request(subagent("child-1", "task-1"));
419 retain(&router, &mut state, &mut first, "model-a").await?;
420
421 let mut second = request(subagent("child-1", "task-2"));
423 let scores = scores(&router, &mut state, &mut second).await?;
424 assert_eq!(
425 scores.first().map(|score| score.target.as_str()),
426 Some("model-a")
427 );
428 Ok(())
429 }
430
431 #[tokio::test]
432 async fn distinct_subagents_are_assigned_independently() -> Result<(), BoxErr> {
433 let router = AffinityRouter::new();
434 let mut state = ();
435
436 retain(
438 &router,
439 &mut state,
440 &mut request(subagent("child-1", "task-1")),
441 "model-a",
442 )
443 .await?;
444
445 let mut sibling = request(subagent("child-2", "task-1"));
447 assert!(scores(&router, &mut state, &mut sibling).await?.is_empty());
448 Ok(())
449 }
450
451 #[tokio::test]
452 async fn subagent_does_not_inherit_session_assignment() -> Result<(), BoxErr> {
453 let router = AffinityRouter::new();
454 let mut state = ();
455
456 retain(
458 &router,
459 &mut state,
460 &mut request(session("session-1", "root-1")),
461 "model-a",
462 )
463 .await?;
464
465 let mut child = request(subagent("child-1", "task-1"));
467 assert!(scores(&router, &mut state, &mut child).await?.is_empty());
468 Ok(())
469 }
470
471 #[tokio::test]
472 async fn classifier_abstains_without_a_session() -> Result<(), BoxErr> {
473 let router = AffinityRouter::new();
474 let mut state = ();
475
476 let mut req = request(Metadata::default());
478 assert!(scores(&router, &mut state, &mut req).await?.is_empty());
479 Ok(())
480 }
481
482 #[tokio::test]
483 async fn message_hash_fallback_uses_the_first_user_message() -> Result<(), BoxErr> {
484 let router = AffinityRouter::new().with_message_hash_fallback();
485 let mut state = ();
486
487 let mut first = task_request(
488 None,
489 "Add a unit test for this function.",
490 Some("Now run the test suite."),
491 );
492 retain(&router, &mut state, &mut first, "weak").await?;
493
494 let mut follow_up = task_request(
495 None,
496 "Add a unit test for this function.",
497 Some("Now file a pull request."),
498 );
499 assert_eq!(
500 scores(&router, &mut state, &mut follow_up)
501 .await?
502 .first()
503 .map(|score| score.target.as_str()),
504 Some("weak")
505 );
506
507 let mut other_task = task_request(
508 None,
509 "Reimplement this binary from two input/output pairs.",
510 Some("Now run the test suite."),
511 );
512 assert!(
513 scores(&router, &mut state, &mut other_task)
514 .await?
515 .is_empty()
516 );
517 Ok(())
518 }
519
520 #[tokio::test]
521 async fn subagents_only_root_traffic_does_not_warn() -> Result<(), BoxErr> {
522 let router = AffinityRouter::for_subagents();
524 let mut state = ();
525
526 let mut root = request(session("session-1", "agent-1"));
527 assert!(scores(&router, &mut state, &mut root).await?.is_empty());
528 assert!(
529 router.should_warn_unkeyed(),
530 "an intentional abstention should leave the warning unconsumed"
531 );
532 Ok(())
533 }
534
535 #[test]
536 fn user_message_hash_ignores_non_text_provider_payloads() {
537 let request = |user_message| Request {
538 llm_request: LlmRequest {
539 messages: vec![user_message],
540 ..LlmRequest::default()
541 },
542 raw_request: None,
543 metadata: None,
544 };
545 let text_only = request(Message::text(Role::User, "Implement the parser."));
546 let text_with_reasoning = request(Message {
547 role: Role::User,
548 content: vec![
549 ContentBlock::Text {
550 text: "Implement the parser.".to_string(),
551 },
552 ContentBlock::Reasoning {
553 text: "Internal provider reasoning.".to_string(),
554 signature: Some("provider-signature".to_string()),
555 details: Vec::new(),
556 },
557 ],
558 });
559
560 assert_eq!(
561 first_user_message_hash(&text_only),
562 first_user_message_hash(&text_with_reasoning)
563 );
564 }
565
566 #[tokio::test]
567 async fn metadata_session_takes_precedence_over_message_hash() -> Result<(), BoxErr> {
568 let router = AffinityRouter::new().with_message_hash_fallback();
569 let mut state = ();
570
571 let mut first = task_request(
572 Some(session("session-1", "agent-a")),
573 "Implement the parser.",
574 None,
575 );
576 retain(&router, &mut state, &mut first, "strong").await?;
577
578 let mut other_session = task_request(
579 Some(session("session-2", "agent-a")),
580 "Implement the parser.",
581 None,
582 );
583 assert!(
584 scores(&router, &mut state, &mut other_session)
585 .await?
586 .is_empty()
587 );
588 Ok(())
589 }
590
591 #[tokio::test]
592 async fn subagent_without_a_session_abstains_and_warns() -> Result<(), BoxErr> {
593 for session_id in [None, Some(String::new())] {
594 let router = AffinityRouter::new().with_message_hash_fallback();
595 let mut state = ();
596 let mut subagent = task_request(
597 Some(Metadata {
598 session_id,
599 agent_id: Some("agent-1".to_string()),
600 is_subagent: true,
601 ..Metadata::default()
602 }),
603 "Implement the parser.",
604 None,
605 );
606
607 retain(&router, &mut state, &mut subagent, "model-a").await?;
608 assert!(scores(&router, &mut state, &mut subagent).await?.is_empty());
609 assert!(
610 !router.should_warn_unkeyed(),
611 "an unidentifiable subagent should consume the warning"
612 );
613 }
614 Ok(())
615 }
616
617 #[tokio::test]
618 async fn one_router_serves_both_roles() -> Result<(), BoxErr> {
619 let router = Arc::new(AffinityRouter::new());
622 let processor: Arc<dyn Processor> = router.clone();
623 let classifier: Arc<dyn Classifier> = router;
624 let mut state = ();
625
626 let mut first = request(session("session-1", "agent-a"));
627 processor
628 .process(
629 &mut state,
630 Event::Decision {
631 request: &mut first,
632 selected_model_id: &fixed_model("model-a"),
633 },
634 )
635 .await?;
636
637 let mut second = request(session("session-1", "agent-b"));
638 let scores = scores(classifier.as_ref(), &mut state, &mut second).await?;
639 assert_eq!(
640 scores.first().map(|score| score.target.as_str()),
641 Some("model-a")
642 );
643 Ok(())
644 }
645
646 #[tokio::test]
647 async fn decision_without_an_affinity_identity_is_ignored() -> Result<(), BoxErr> {
648 let router = AffinityRouter::new();
649 let mut state = ();
650 let mut unkeyed = request(Metadata::default());
651
652 router
653 .process(
654 &mut state,
655 Event::Decision {
656 request: &mut unkeyed,
657 selected_model_id: &fixed_model("model-a"),
658 },
659 )
660 .await?;
661
662 let mut req = request(session("session-1", "agent-a"));
663 assert!(scores(&router, &mut state, &mut req).await?.is_empty());
664 Ok(())
665 }
666
667 #[tokio::test]
668 async fn decisions_retain_their_originating_request_identity() -> Result<(), BoxErr> {
669 let router = AffinityRouter::new();
670 let mut state = ();
671 let mut first = request(session("session-1", "agent-a"));
672 let mut second = request(session("session-2", "agent-b"));
673
674 router
677 .process(
678 &mut state,
679 Event::Decision {
680 request: &mut second,
681 selected_model_id: &fixed_model("model-b"),
682 },
683 )
684 .await?;
685 router
686 .process(
687 &mut state,
688 Event::Decision {
689 request: &mut first,
690 selected_model_id: &fixed_model("model-a"),
691 },
692 )
693 .await?;
694
695 let first_scores = scores(&router, &mut state, &mut first).await?;
696 let second_scores = scores(&router, &mut state, &mut second).await?;
697 assert_eq!(
698 first_scores.first().map(|score| score.target.as_str()),
699 Some("model-a")
700 );
701 assert_eq!(
702 second_scores.first().map(|score| score.target.as_str()),
703 Some("model-b")
704 );
705 Ok(())
706 }
707
708 #[tokio::test]
709 async fn distinct_sessions_are_assigned_independently() -> Result<(), BoxErr> {
710 let router = AffinityRouter::new();
711 let mut state = ();
712
713 retain(
714 &router,
715 &mut state,
716 &mut request(session("session-1", "agent-a")),
717 "model-a",
718 )
719 .await?;
720 retain(
721 &router,
722 &mut state,
723 &mut request(session("session-2", "agent-a")),
724 "model-b",
725 )
726 .await?;
727
728 let first = scores(
729 &router,
730 &mut state,
731 &mut request(session("session-1", "other")),
732 )
733 .await?;
734 let second = scores(
735 &router,
736 &mut state,
737 &mut request(session("session-2", "other")),
738 )
739 .await?;
740 assert_eq!(
741 first.first().map(|score| score.target.as_str()),
742 Some("model-a")
743 );
744 assert_eq!(
745 second.first().map(|score| score.target.as_str()),
746 Some("model-b")
747 );
748 Ok(())
749 }
750
751 #[tokio::test]
752 async fn subagent_without_an_agent_id_abstains_and_warns() -> Result<(), BoxErr> {
753 let router = AffinityRouter::new();
754 let mut state = ();
755
756 let metadata = Metadata {
759 session_id: Some("session-1".to_string()),
760 is_subagent: true,
761 ..Metadata::default()
762 };
763 let mut req = request(metadata);
764 retain(&router, &mut state, &mut req, "model-a").await?;
765 assert!(scores(&router, &mut state, &mut req).await?.is_empty());
766 assert!(
767 !router.should_warn_unkeyed(),
768 "an unidentifiable subagent should consume the warning"
769 );
770 Ok(())
771 }
772
773 #[tokio::test]
774 async fn assignments_are_bounded_by_the_cap() -> Result<(), BoxErr> {
775 let router = AffinityRouter::new();
776 let mut state = ();
777
778 for index in 0..=MAX_ASSIGNMENTS {
780 let session_id = format!("session-{index}");
781 retain(
782 &router,
783 &mut state,
784 &mut request(session(&session_id, "agent-a")),
785 "model-a",
786 )
787 .await?;
788 }
789
790 let len = router.assignments.lock().len();
791 assert_eq!(len, MAX_ASSIGNMENTS);
792 Ok(())
793 }
794
795 #[tokio::test]
796 async fn release_on_user_turn_drops_the_assignment_only_when_the_user_speaks()
797 -> Result<(), BoxErr> {
798 let router = AffinityRouter::new().with_release_on_user_turn();
799 let mut state = ();
800 let mut opening = task_request(Some(session("session-1", "agent-a")), "add caching", None);
801 retain(&router, &mut state, &mut opening, "weak").await?;
802
803 let mut continued = opening.clone();
805 continued.llm_request.messages.push(Message {
806 role: Role::User,
807 content: vec![ContentBlock::ToolResult(ToolResult {
808 tool_call_id: "call-1".to_string(),
809 content: Vec::new(),
810 is_error: None,
811 })],
812 });
813 assert_eq!(
814 scores(&router, &mut state, &mut continued)
815 .await?
816 .first()
817 .map(|s| s.target.as_str()),
818 Some("weak")
819 );
820
821 let mut spoke = task_request(
823 Some(session("session-1", "agent-a")),
824 "add caching",
825 Some("no, shared across processes"),
826 );
827 assert!(scores(&router, &mut state, &mut spoke).await?.is_empty());
828 Ok(())
829 }
830
831 #[tokio::test]
832 async fn latch_only_retains_matching_models() -> Result<(), BoxErr> {
833 let router = AffinityRouter::new().with_latch_only(["strong"]);
834 let mut state = ();
835 let mut req = request(session("session-1", "agent-a"));
836
837 retain(&router, &mut state, &mut req, "weak").await?;
839 assert!(scores(&router, &mut state, &mut req).await?.is_empty());
840
841 retain(&router, &mut state, &mut req, "strong").await?;
843 assert_eq!(
844 scores(&router, &mut state, &mut req)
845 .await?
846 .first()
847 .map(|s| s.target.as_str()),
848 Some("strong")
849 );
850 Ok(())
851 }
852}