Skip to main content

theway_core/agent/runtime_extensions/
context.rs

1use 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    /// Replace the live branch projection while retaining all shared handles.
72    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    /// Add the de-duplicated model-visible projection to one normalized model
82    /// request. This never mutates the persisted agent transcript.
83    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    /// Project the de-duplicated model-visible entries into one compaction-only
116    /// message list. Private state and custom events never enter this list.
117    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}