theway_core/agent/runtime_extensions/
context.rs1use std::collections::BTreeMap;
2use std::sync::Arc;
3
4use parking_lot::RwLock;
5use serde_json::Value;
6use theway_contract::extension::{
7 ExtensionDurableEntry, ExtensionDurableEntryPayload, ExtensionModelContextPlacement,
8};
9use thiserror::Error;
10
11use crate::types::{AgentMessage, CustomMessage};
12
13#[derive(Clone, Debug, PartialEq, Eq)]
14pub struct ExtensionModelContextItem {
15 pub extension_id: String,
16 pub context_id: String,
17 pub placement: ExtensionModelContextPlacement,
18 pub content: Value,
19}
20
21#[derive(Clone, Debug, Default)]
22pub struct ExtensionModelContextProjection {
23 items: Arc<RwLock<Vec<ExtensionModelContextItem>>>,
24}
25
26impl ExtensionModelContextProjection {
27 pub fn rebuild(
28 entries: impl IntoIterator<Item = ExtensionDurableEntry>,
29 ) -> Result<Self, ExtensionModelContextProjectionError> {
30 let mut positions = BTreeMap::<(String, String), usize>::new();
31 let mut items = Vec::new();
32 for entry in entries {
33 entry.validate().map_err(|error| {
34 ExtensionModelContextProjectionError::InvalidEntry(error.to_string())
35 })?;
36 let ExtensionDurableEntryPayload::ModelContext {
37 context_id,
38 placement,
39 content,
40 } = entry.entry
41 else {
42 continue;
43 };
44 let key = (entry.extension_id.clone(), context_id.clone());
45 let projected = ExtensionModelContextItem {
46 extension_id: entry.extension_id,
47 context_id,
48 placement,
49 content,
50 };
51 if let Some(index) = positions.get(&key).copied() {
52 items[index] = projected;
53 } else {
54 positions.insert(key, items.len());
55 items.push(projected);
56 }
57 }
58 Ok(Self {
59 items: Arc::new(RwLock::new(items)),
60 })
61 }
62
63 pub fn items(&self) -> Vec<ExtensionModelContextItem> {
64 self.items.read().clone()
65 }
66
67 pub fn into_items(self) -> Vec<ExtensionModelContextItem> {
68 self.items.read().clone()
69 }
70
71 pub fn replace(
73 &self,
74 entries: impl IntoIterator<Item = ExtensionDurableEntry>,
75 ) -> Result<(), ExtensionModelContextProjectionError> {
76 let rebuilt = Self::rebuild(entries)?;
77 *self.items.write() = rebuilt.items();
78 Ok(())
79 }
80
81 pub fn apply_to_request(
84 &self,
85 request: &mut crate::agent::model_request::NormalizedModelRequestDraft,
86 ) {
87 let items = self.items.read();
88 let sections = items
89 .iter()
90 .filter_map(|item| {
91 (item.placement == ExtensionModelContextPlacement::SystemPromptSection)
92 .then(|| item.content.as_str())
93 .flatten()
94 })
95 .collect::<Vec<_>>();
96 if !sections.is_empty() {
97 let suffix = sections.join("\n\n");
98 request.system_instructions = Some(match request.system_instructions.take() {
99 Some(base) if !base.is_empty() => format!("{base}\n\n{suffix}"),
100 _ => suffix,
101 });
102 }
103 request.messages.extend(items.iter().filter_map(|item| {
104 if item.placement != ExtensionModelContextPlacement::Message {
105 return None;
106 }
107 let message = serde_json::from_value::<AgentMessage>(item.content.clone()).ok()?;
108 match message {
109 AgentMessage::Llm(message) => Some(message),
110 AgentMessage::Custom(_) => None,
111 }
112 }));
113 }
114
115 pub fn compaction_messages(&self) -> Vec<AgentMessage> {
118 self.items
119 .read()
120 .iter()
121 .map(|item| match item.placement {
122 ExtensionModelContextPlacement::Message => {
123 serde_json::from_value(item.content.clone())
124 .unwrap_or_else(|_| model_context_marker(item, "message"))
125 }
126 ExtensionModelContextPlacement::SystemPromptSection => {
127 model_context_marker(item, "system_prompt_section")
128 }
129 })
130 .collect()
131 }
132}
133
134impl PartialEq for ExtensionModelContextProjection {
135 fn eq(&self, other: &Self) -> bool {
136 *self.items.read() == *other.items.read()
137 }
138}
139
140impl Eq for ExtensionModelContextProjection {}
141
142fn model_context_marker(item: &ExtensionModelContextItem, placement: &str) -> AgentMessage {
143 AgentMessage::Custom(CustomMessage {
144 role: "extension_model_context".into(),
145 timestamp: 0,
146 payload: serde_json::json!({
147 "extensionId": item.extension_id,
148 "contextId": item.context_id,
149 "placement": placement,
150 "content": item.content,
151 }),
152 })
153}
154
155#[derive(Clone, Debug, Error, PartialEq, Eq)]
156pub enum ExtensionModelContextProjectionError {
157 #[error("persistent model-context entry is invalid: {0}")]
158 InvalidEntry(String),
159}