21
use std::sync::Arc;
22
23
use async_trait::async_trait;
24
-use onde::inference::{ChatEngine, ToolDefinition};
24
+use onde::inference::{ChatEngine, ChatMessage, ChatRole, ToolDefinition};
25
use serde::Deserialize;
26
use tokio::sync::Mutex;
27
62
/// Backend errors are plain strings. Callers map them to ACP errors.
63
pub type BackendError = String;
64
65
+/// Rough context budget for a conversation, in estimated tokens (see
66
+/// [`estimate_tokens`]). When a snapshot exceeds this, the agent loops compact
67
+/// history before the next tool round.
68
+pub const DEFAULT_CONTEXT_TOKEN_BUDGET: usize = 24_000;
69
+
70
+/// How many trailing messages survive a compaction verbatim (the rest are
71
+/// folded into the summary).
72
+pub const COMPACT_KEEP_LAST: usize = 6;
73
+
74
+/// The summarization request sent to the model when compacting history.
75
+const SUMMARIZE_PROMPT: &str = "Summarize this coding session so far: decisions made, \
76
+ files touched, current state, open items. Be concise and factual.";
77
+
78
+/// Crude token estimate for a history snapshot: serialized characters / 4.
79
+/// Deliberately model-agnostic — it only needs to be in the right ballpark to
80
+/// decide when compaction is worth an extra inference round.
81
+pub fn estimate_tokens(history: &[serde_json::Value]) -> usize {
82
+ let chars: usize = history
83
+ .iter()
84
+ .map(|message| message.to_string().chars().count())
85
+ .sum();
86
+ chars / 4
87
+}
88
+
89
/// A sink for streaming assistant text deltas to the UI as they are produced.
90
///
91
/// When a caller passes `Some(sink)`, a streaming-capable backend forwards each
134
/// than on-device. Drives UI labelling so the displayed model can't claim a
135
/// local model while requests actually go to the cloud.
136
fn is_remote(&self) -> bool;
137
+
138
+ /// A serializable snapshot of the conversation history, one JSON object per
139
+ /// message (`{"role": ..., "content": ...}` at minimum). The snapshot is
140
+ /// what the session store persists; it includes any seeded system message
141
+ /// so [`InferenceBackend::restore_history`] can replace state wholesale.
142
+ async fn history_snapshot(&self) -> Vec<serde_json::Value>;
143
+
144
+ /// Replace the conversation history with a previously saved snapshot.
145
+ /// Backends that cannot represent every entry (e.g. on-device history has
146
+ /// no tool-call structure) flatten what they can and drop the rest.
147
+ async fn restore_history(&self, history: Vec<serde_json::Value>);
148
+
149
+ /// Shrink the conversation history: summarize everything so far with one
150
+ /// extra (non-streaming) inference round, then rebuild history as
151
+ /// `[system message, summary, last keep_last non-system messages]`. On
152
+ /// error the original history is left in place.
153
+ async fn compact_history(&self, keep_last: usize) -> Result<(), BackendError>;
154
}
155
156
// ── Local backend (onde ChatEngine) ──────────────────────────────────────────────
254
fn is_remote(&self) -> bool {
255
false
256
}
257
+
258
+ async fn history_snapshot(&self) -> Vec<serde_json::Value> {
259
+ // onde's `history()` already flattens tool entries: assistant tool
260
+ // calls become plain assistant text and tool results are omitted, so
261
+ // the snapshot is lossy for tool-heavy turns (acceptable in this MVP).
262
+ self.engine
263
+ .history()
264
+ .await
265
+ .iter()
266
+ .map(|message| {
267
+ serde_json::json!({
268
+ "role": message.role.to_string(),
269
+ "content": message.content,
270
+ })
271
+ })
272
+ .collect()
273
+ }
274
+
275
+ async fn restore_history(&self, history: Vec<serde_json::Value>) {
276
+ self.engine.clear_history().await;
277
+ for entry in history {
278
+ let role = entry["role"].as_str().unwrap_or("");
279
+ let content = entry["content"].as_str().unwrap_or("").to_string();
280
+ // Tool-call-only assistant entries and empty tool results carry no
281
+ // text a plain chat history can replay; drop them.
282
+ if content.is_empty() && role != "user" && role != "system" {
283
+ continue;
284
+ }
285
+ let message = match role {
286
+ "system" => ChatMessage::system(content),
287
+ "user" => ChatMessage::user(content),
288
+ "assistant" => ChatMessage::assistant(content),
289
+ // Tool results flatten to plain text (MVP; acceptable loss).
290
+ "tool" => ChatMessage::user(format!("[tool result]\n{content}")),
291
+ _ => continue,
292
+ };
293
+ self.engine.push_history(message).await;
294
+ }
295
+ }
296
+
297
+ async fn compact_history(&self, keep_last: usize) -> Result<(), BackendError> {
298
+ let snapshot = self.engine.history().await;
299
+ // One plain (tool-free) inference round produces the summary. On error
300
+ // history is untouched — send_message only mutates it on success, and
301
+ // whatever it appended is wiped by the clear below anyway.
302
+ let result = self
303
+ .engine
304
+ .send_message(SUMMARIZE_PROMPT)
305
+ .await
306
+ .map_err(|error| error.to_string())?;
307
+ // Local models may reason in <think> blocks; keep only the visible part.
308
+ let (_think, summary) = crate::chat::strip_think_blocks(&result.text);
309
+
310
+ self.engine.clear_history().await;
311
+ // Leading system messages carry the session context; keep them all.
312
+ for message in snapshot
313
+ .iter()
314
+ .take_while(|message| message.role == ChatRole::System)
315
+ {
316
+ self.engine.push_history(message.clone()).await;
317
+ }
318
+ self.engine
319
+ .push_history(ChatMessage::user(format!(
320
+ "[Conversation summary]\n{summary}"
321
+ )))
322
+ .await;
323
+ let non_system: Vec<&ChatMessage> = snapshot
324
+ .iter()
325
+ .filter(|message| message.role != ChatRole::System)
326
+ .collect();
327
+ let tail_start = non_system.len().saturating_sub(keep_last);
328
+ for message in &non_system[tail_start..] {
329
+ self.engine.push_history((*message).clone()).await;
330
+ }
331
+ Ok(())
332
+ }
333
}
334
335
/// Drain an onde streaming receiver, forwarding each token to `sink` and
708
fn is_remote(&self) -> bool {
709
true
710
}
711
+
712
+ async fn history_snapshot(&self) -> Vec<serde_json::Value> {
713
+ self.history.lock().await.clone()
714
+ }
715
+
716
+ async fn restore_history(&self, history: Vec<serde_json::Value>) {
717
+ // The snapshot includes the seeded system message, so a wholesale
718
+ // replacement restores exactly what was saved.
719
+ *self.history.lock().await = history;
720
+ }
721
+
722
+ async fn compact_history(&self, keep_last: usize) -> Result<(), BackendError> {
723
+ let snapshot: Vec<serde_json::Value> = self.history.lock().await.clone();
724
+
725
+ // Ask the endpoint for a summary of the conversation so far, through
726
+ // the ordinary completion machinery (non-streaming).
727
+ self.history
728
+ .lock()
729
+ .await
730
+ .push(serde_json::json!({ "role": "user", "content": SUMMARIZE_PROMPT }));
731
+ let summary = match self.complete(None, None).await {
732
+ Ok(result) => result.text,
733
+ Err(error) => {
734
+ // Roll back the summarization request; the turn never happened.
735
+ *self.history.lock().await = snapshot;
736
+ return Err(error);
737
+ }
738
+ };
739
+
740
+ let system = snapshot
741
+ .first()
742
+ .filter(|message| message["role"] == "system")
743
+ .cloned();
744
+ let non_system: Vec<serde_json::Value> = snapshot
745
+ .iter()
746
+ .filter(|message| message["role"] != "system")
747
+ .cloned()
748
+ .collect();
749
+ let tail_start = non_system.len().saturating_sub(keep_last);
750
+ let mut tail = non_system[tail_start..].to_vec();
751
+ // Drop leading tool results whose assistant tool-call message was
752
+ // summarized away — strict endpoints reject orphaned `role: "tool"`
753
+ // entries on the very next request.
754
+ while tail
755
+ .first()
756
+ .is_some_and(|message| message["role"] == "tool")
757
+ {
758
+ tail.remove(0);
759
+ }
760
+
761
+ let mut rebuilt = Vec::new();
762
+ if let Some(system) = system {
763
+ rebuilt.push(system);
764
+ }
765
+ rebuilt.push(serde_json::json!({
766
+ "role": "user",
767
+ "content": format!("[Conversation summary]\n{summary}"),
768
+ }));
769
+ rebuilt.extend(tail);
770
+ *self.history.lock().await = rebuilt;
771
+ Ok(())
772
+ }
773
}
774
775
// ── OpenAI response shapes ────────────────────────────────────────────────────────
962
assert_eq!(last["content"], "cancelled by the user");
963
}
964
965
+ #[test]
966
+ fn estimate_tokens_scales_with_serialized_size() {
967
+ assert_eq!(estimate_tokens(&[]), 0);
968
+
969
+ let short = vec![serde_json::json!({ "role": "user", "content": "hi" })];
970
+ let long = vec![serde_json::json!({ "role": "user", "content": "x".repeat(4_000) })];
971
+ let short_estimate = estimate_tokens(&short);
972
+ let long_estimate = estimate_tokens(&long);
973
+
974
+ assert!(short_estimate > 0, "non-empty history estimates > 0 tokens");
975
+ assert!(long_estimate > short_estimate, "longer history costs more");
976
+ // 4,000 content chars / 4 ≈ 1,000 tokens, plus a little JSON framing.
977
+ assert!((1_000..1_100).contains(&long_estimate), "{long_estimate}");
978
+ }
979
+
980
+ #[tokio::test]
981
+ async fn openai_snapshot_restore_round_trips_exactly() {
982
+ let backend = OpenAiBackend::new("http://localhost", "", "m", Some("be helpful".into()));
983
+ {
984
+ let mut history = backend.history.lock().await;
985
+ history.push(serde_json::json!({ "role": "user", "content": "hello" }));
986
+ history.push(streamed_assistant_history(
987
+ "",
988
+ &[ToolCall {
989
+ id: "call_1".to_string(),
990
+ name: "read_file".to_string(),
991
+ arguments: r#"{"path":"a.rs"}"#.to_string(),
992
+ }],
993
+ ));
994
+ history.push(serde_json::json!({
995
+ "role": "tool", "tool_call_id": "call_1", "content": "fn main() {}",
996
+ }));
997
+ history.push(serde_json::json!({ "role": "assistant", "content": "done" }));
998
+ }
999
+ let snapshot = backend.history_snapshot().await;
1000
+ assert_eq!(
1001
+ snapshot[0]["role"], "system",
1002
+ "snapshot keeps the system message"
1003
+ );
1004
+
1005
+ // Restoring into a backend seeded with a *different* system prompt must
1006
+ // replace everything, including that seed.
1007
+ let restored = OpenAiBackend::new("http://localhost", "", "m", Some("other seed".into()));
1008
+ restored.restore_history(snapshot.clone()).await;
1009
+ assert_eq!(restored.history_snapshot().await, snapshot);
1010
+ }
1011
+
1012
+ /// Minimal scripted OpenAI-compatible endpoint: accepts one HTTP request on
1013
+ /// a std listener and answers with a fixed non-streaming completion.
1014
+ fn spawn_completion_stub(summary: &str) -> std::net::SocketAddr {
1015
+ use std::io::{Read, Write};
1016
+
1017
+ let listener = std::net::TcpListener::bind("127.0.0.1:0").unwrap();
1018
+ let addr = listener.local_addr().unwrap();
1019
+ let body = serde_json::json!({
1020
+ "choices": [{ "message": { "role": "assistant", "content": summary } }]
1021
+ })
1022
+ .to_string();
1023
+ std::thread::spawn(move || {
1024
+ let (mut stream, _) = listener.accept().unwrap();
1025
+ // Read until the full request (headers + content-length body) is in.
1026
+ let mut request = Vec::new();
1027
+ let mut chunk = [0u8; 4096];
1028
+ loop {
1029
+ let n = stream.read(&mut chunk).unwrap_or(0);
1030
+ if n == 0 {
1031
+ break;
1032
+ }
1033
+ request.extend_from_slice(&chunk[..n]);
1034
+ if let Some(headers_end) =
1035
+ request.windows(4).position(|window| window == b"\r\n\r\n")
1036
+ {
1037
+ let headers = String::from_utf8_lossy(&request[..headers_end]);
1038
+ let content_length = headers
1039
+ .lines()
1040
+ .find_map(|line| {
1041
+ line.to_ascii_lowercase()
1042
+ .strip_prefix("content-length:")
1043
+ .map(|value| value.trim().parse::<usize>().unwrap_or(0))
1044
+ })
1045
+ .unwrap_or(0);
1046
+ if request.len() >= headers_end + 4 + content_length {
1047
+ break;
1048
+ }
1049
+ }
1050
+ }
1051
+ let response = format!(
1052
+ "HTTP/1.1 200 OK\r\ncontent-type: application/json\r\n\
1053
+ content-length: {}\r\nconnection: close\r\n\r\n{}",
1054
+ body.len(),
1055
+ body
1056
+ );
1057
+ let _ = stream.write_all(response.as_bytes());
1058
+ });
1059
+ addr
1060
+ }
1061
+
1062
+ #[tokio::test]
1063
+ async fn compact_history_rebuilds_system_summary_and_tail() {
1064
+ let addr = spawn_completion_stub("We refactored backend.rs; tests pass.");
1065
+ let backend = OpenAiBackend::new(
1066
+ format!("http://{addr}/v1"),
1067
+ "test-key",
1068
+ "test-model",
1069
+ Some("be helpful".into()),
1070
+ );
1071
+ {
1072
+ let mut history = backend.history.lock().await;
1073
+ for i in 0..5 {
1074
+ let role = if i % 2 == 0 { "user" } else { "assistant" };
1075
+ history.push(serde_json::json!({
1076
+ "role": role, "content": format!("message {i}"),
1077
+ }));
1078
+ }
1079
+ }
1080
+
1081
+ backend.compact_history(2).await.unwrap();
1082
+
1083
+ let history = backend.history_snapshot().await;
1084
+ assert_eq!(history.len(), 4, "system + summary + last 2: {history:?}");
1085
+ assert_eq!(history[0]["role"], "system");
1086
+ assert_eq!(history[0]["content"], "be helpful");
1087
+ assert_eq!(history[1]["role"], "user");
1088
+ let summary_text = history[1]["content"].as_str().unwrap();
1089
+ assert!(summary_text.starts_with("[Conversation summary]\n"));
1090
+ assert!(summary_text.contains("We refactored backend.rs; tests pass."));
1091
+ assert_eq!(
1092
+ history[2],
1093
+ serde_json::json!({ "role": "assistant", "content": "message 3" })
1094
+ );
1095
+ assert_eq!(
1096
+ history[3],
1097
+ serde_json::json!({ "role": "user", "content": "message 4" })
1098
+ );
1099
+ }
1100
+
1101
+ #[tokio::test]
1102
+ async fn compact_history_failure_leaves_history_intact() {
1103
+ // No listener at this address: the summarization request fails, and
1104
+ // history must roll back to exactly what it was.
1105
+ let backend =
1106
+ OpenAiBackend::new("http://127.0.0.1:9", "", "test-model", Some("sys".into()));
1107
+ backend
1108
+ .history
1109
+ .lock()
1110
+ .await
1111
+ .push(serde_json::json!({ "role": "user", "content": "hello" }));
1112
+ let before = backend.history_snapshot().await;
1113
+
1114
+ assert!(backend.compact_history(2).await.is_err());
1115
+ assert_eq!(backend.history_snapshot().await, before);
1116
+ }
1117
+
1118
#[test]
1119
fn assistant_message_with_tool_calls_round_trips() {
1120
let message = ResponseMessage {