oxicode_agent/advisor/
agent_advisor.rs1use std::sync::Arc;
18
19use async_trait::async_trait;
20
21use crate::Agent;
22use crate::advisor::runtime::AdvisorAgent;
23
24pub type AdvisorPromptHook = Arc<dyn Fn(&Agent) + Send + Sync>;
27
28pub struct AgentAdvisor {
30 agent: Arc<Agent>,
31 on_prompted: Option<AdvisorPromptHook>,
32}
33
34impl AgentAdvisor {
35 #[must_use]
37 pub fn new(agent: Arc<Agent>) -> Self {
38 Self {
39 agent,
40 on_prompted: None,
41 }
42 }
43
44 #[must_use]
48 pub fn with_post_prompt_hook(agent: Arc<Agent>, hook: AdvisorPromptHook) -> Self {
49 Self {
50 agent,
51 on_prompted: Some(hook),
52 }
53 }
54
55 #[must_use]
57 pub fn agent(&self) -> &Agent {
58 &self.agent
59 }
60
61 #[must_use]
63 pub fn into_agent(self) -> Arc<Agent> {
64 self.agent
65 }
66}
67
68#[async_trait]
69impl AdvisorAgent for AgentAdvisor {
70 async fn prompt(&self, input: String) -> Result<(), String> {
71 self.agent
76 .continue_with(input)
77 .await
78 .map(|_| {
79 if let Some(hook) = &self.on_prompted {
80 hook(&self.agent);
81 }
82 })
83 .map_err(|e| e.to_string())
84 }
85
86 fn abort(&self, _reason: &str) {
87 self.agent.cancel();
90 }
91
92 fn reset(&self) {
93 self.agent.reset();
94 }
95
96 async fn rollback_to(&self, count: usize) {
97 self.agent.update_state(|s| s.messages.truncate(count));
98 }
99
100 fn message_count(&self) -> usize {
101 self.agent.state().messages.len()
102 }
103}
104
105#[cfg(test)]
106mod tests {
107 #![allow(clippy::unwrap_used)]
108 use super::*;
109 use crate::config::AgentConfig;
110 use oxicode_ai::{Message, Provider};
111
112 struct NopProvider;
115 impl Provider for NopProvider {
116 fn stream<'a>(
117 &'a self,
118 _model: &'a oxicode_ai::Model,
119 _context: &'a oxicode_ai::Context,
120 _options: Option<oxicode_ai::StreamOptions>,
121 ) -> std::pin::Pin<
122 Box<dyn std::future::Future<Output = oxicode_ai::StreamResult> + Send + 'a>,
123 > {
124 let s: std::pin::Pin<
127 Box<dyn futures::Stream<Item = oxicode_ai::ProviderEvent> + Send>,
128 > = Box::pin(futures::stream::empty::<oxicode_ai::ProviderEvent>());
129 Box::pin(async move { Ok(s) })
130 }
131 }
132
133 #[tokio::test]
134 async fn message_count_tracks_state() {
135 let provider: Arc<dyn Provider> = Arc::new(NopProvider);
136 let agent = Arc::new(Agent::new_empty(provider, AgentConfig::default()));
137 let advisor = AgentAdvisor::new(Arc::clone(&agent));
138
139 assert_eq!(advisor.message_count(), 0);
140 agent.update_state(|s| {
142 s.messages.push(Message::user("hello"));
143 s.messages.push(Message::user("world"));
144 });
145 assert_eq!(advisor.message_count(), 2);
146 }
147
148 #[tokio::test]
149 async fn rollback_to_truncates_messages() {
150 let provider: Arc<dyn Provider> = Arc::new(NopProvider);
151 let agent = Arc::new(Agent::new_empty(provider, AgentConfig::default()));
152 let advisor = AgentAdvisor::new(Arc::clone(&agent));
153 agent.update_state(|s| {
154 s.messages.push(Message::user("a"));
155 s.messages.push(Message::user("b"));
156 s.messages.push(Message::user("c"));
157 s.messages.push(Message::user("d"));
158 });
159 advisor.rollback_to(2).await;
160 assert_eq!(advisor.message_count(), 2);
161 assert_eq!(agent.state().messages[0].text_content().unwrap(), "a");
162 }
163
164 #[tokio::test]
165 async fn reset_clears_state() {
166 let provider: Arc<dyn Provider> = Arc::new(NopProvider);
167 let agent = Arc::new(Agent::new_empty(provider, AgentConfig::default()));
168 let advisor = AgentAdvisor::new(Arc::clone(&agent));
169 agent.update_state(|s| {
170 s.messages.push(Message::user("a"));
171 });
172 assert_eq!(advisor.message_count(), 1);
173 advisor.reset();
174 assert_eq!(advisor.message_count(), 0);
175 }
176
177 #[test]
178 fn agent_accessor_and_into_agent_round_trip() {
179 let provider: Arc<dyn Provider> = Arc::new(NopProvider);
180 let agent = Arc::new(Agent::new_empty(provider, AgentConfig::default()));
181 let cloned = Arc::clone(&agent);
182 let advisor = AgentAdvisor::new(cloned);
183 assert!(std::ptr::eq(advisor.agent(), Arc::as_ref(&agent)));
185 assert!(Arc::ptr_eq(&advisor.into_agent(), &agent));
186 }
187}