1use serde::{Deserialize, Serialize};
9use surrealdb::Surreal;
10
11use super::error::GraphError;
12use super::store::Db;
13
14#[derive(Debug, Clone, Copy, PartialEq, Serialize, Deserialize)]
16#[serde(rename_all = "snake_case")]
17pub enum OutcomeKind {
18 Success,
19 Partial,
20 Failed,
21}
22
23impl OutcomeKind {
24 #[must_use]
26 pub fn reward(self) -> f64 {
27 match self {
28 Self::Success => 1.0,
29 Self::Partial => 0.5,
30 Self::Failed => 0.0,
31 }
32 }
33}
34
35impl std::fmt::Display for OutcomeKind {
36 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
37 match self {
38 Self::Success => write!(f, "success"),
39 Self::Partial => write!(f, "partial"),
40 Self::Failed => write!(f, "failed"),
41 }
42 }
43}
44
45impl std::str::FromStr for OutcomeKind {
46 type Err = String;
47
48 fn from_str(s: &str) -> Result<Self, Self::Err> {
49 match s.to_lowercase().as_str() {
50 "success" => Ok(Self::Success),
51 "partial" => Ok(Self::Partial),
52 "failed" => Ok(Self::Failed),
53 other => Err(format!("unknown outcome kind: {other}")),
54 }
55 }
56}
57
58pub const DEFAULT_UTILITY: f64 = 0.5;
60
61const USED_ALPHA: f64 = 0.1;
63
64const UNUSED_ALPHA: f64 = 0.05;
66
67const UNUSED_REWARD: f64 = 0.3;
69
70#[derive(Debug, Clone, Default)]
72pub struct FeedbackReport {
73 pub outcome_entity_id: String,
74 pub edges_created: u32,
75 pub entities_updated: u32,
76 pub errors: Vec<String>,
77}
78
79pub async fn record_outcome_feedback(
81 db: &Surreal<Db>,
82 session_id: &str,
83 outcome: OutcomeKind,
84 retrieved_entity_ids: &[String],
85 used_entity_ids: Option<&[String]>,
86) -> Result<FeedbackReport, GraphError> {
87 let mut report = FeedbackReport::default();
88
89 if retrieved_entity_ids.is_empty() {
90 return Ok(report);
91 }
92
93 let outcome_id = create_outcome_entity(db, session_id, outcome).await?;
94 report.outcome_entity_id = outcome_id.clone();
95
96 let reward = outcome.reward();
97
98 let used_set: Option<std::collections::HashSet<&str>> =
100 used_entity_ids.map(|ids| ids.iter().map(|s| s.as_str()).collect());
101
102 let outcome_id_ref = &outcome_id;
104 let futures: Vec<_> = retrieved_entity_ids
105 .iter()
106 .map(|entity_id| {
107 let was_used = used_set
108 .as_ref()
109 .map(|s| s.contains(entity_id.as_str()))
110 .unwrap_or(true);
111 let (alpha, effective_reward) = if was_used {
112 (USED_ALPHA, reward)
113 } else {
114 (UNUSED_ALPHA, UNUSED_REWARD)
115 };
116
117 async move {
118 let edge_result = create_contribution_edge(
119 db,
120 entity_id,
121 outcome_id_ref,
122 outcome,
123 was_used,
124 session_id,
125 )
126 .await;
127 let utility_result =
128 update_utility_score(db, entity_id, effective_reward, alpha).await;
129 (entity_id, edge_result, utility_result)
130 }
131 })
132 .collect();
133
134 let results = futures::future::join_all(futures).await;
135
136 for (entity_id, edge_result, utility_result) in results {
137 match edge_result {
138 Ok(()) => report.edges_created += 1,
139 Err(e) => {
140 report
141 .errors
142 .push(format!("edge {entity_id} -> {outcome_id}: {e}"));
143 }
144 }
145 match utility_result {
146 Ok(()) => report.entities_updated += 1,
147 Err(e) => {
148 report
149 .errors
150 .push(format!("utility update {entity_id}: {e}"));
151 }
152 }
153 }
154
155 Ok(report)
156}
157
158async fn create_outcome_entity(
159 db: &Surreal<Db>,
160 session_id: &str,
161 outcome: OutcomeKind,
162) -> Result<String, GraphError> {
163 let abstract_text = format!("Session {session_id} outcome: {outcome}");
164
165 let mut response = db
166 .query(
167 r#"
168 CREATE entity SET
169 name = $name,
170 entity_type = "outcome",
171 abstract = $abstract,
172 overview = "",
173 content = NONE,
174 attributes = $attributes,
175 embedding = NONE,
176 mutable = false,
177 access_count = 0,
178 utility_score = $utility,
179 utility_updates = 0,
180 created_at = time::now(),
181 updated_at = time::now(),
182 source = $source
183 "#,
184 )
185 .bind(("name", format!("outcome-{session_id}")))
186 .bind(("abstract", abstract_text))
187 .bind((
188 "attributes",
189 serde_json::json!({
190 "outcome_result": outcome.to_string(),
191 "session_id": session_id,
192 }),
193 ))
194 .bind(("utility", DEFAULT_UTILITY))
195 .bind(("source", format!("caliber:{session_id}")))
196 .await?;
197
198 let entity: Option<super::types::Entity> = super::deserialize_take_opt(&mut response, 0)?;
199 let entity = entity.ok_or_else(|| {
200 GraphError::Db(surrealdb::Error::thrown(
201 "failed to create outcome entity".into(),
202 ))
203 })?;
204
205 Ok(entity.id_string())
206}
207
208async fn create_contribution_edge(
209 db: &Surreal<Db>,
210 entity_id: &str,
211 outcome_id: &str,
212 outcome: OutcomeKind,
213 was_used: bool,
214 session_id: &str,
215) -> Result<(), GraphError> {
216 db.query(
217 r#"
218 LET $from = type::record($from_id);
219 LET $to = type::record($to_id);
220 RELATE $from -> contributed_to -> $to SET
221 outcome_result = $outcome_result,
222 was_used = $was_used,
223 session_id = $session_id,
224 timestamp = time::now()
225 "#,
226 )
227 .bind(("from_id", entity_id.to_string()))
228 .bind(("to_id", outcome_id.to_string()))
229 .bind(("outcome_result", outcome.to_string()))
230 .bind(("was_used", was_used))
231 .bind(("session_id", session_id.to_string()))
232 .await?
233 .check()?;
234
235 Ok(())
236}
237
238async fn update_utility_score(
240 db: &Surreal<Db>,
241 entity_id: &str,
242 reward: f64,
243 alpha: f64,
244) -> Result<(), GraphError> {
245 db.query(
248 r#"
249 LET $raw = (1.0 - $alpha) * type::record($id).utility_score + $alpha * $reward;
250 LET $clamped = IF $raw < 0.0 THEN 0.0 ELSE IF $raw > 1.0 THEN 1.0 ELSE $raw END END;
251 UPDATE type::record($id) SET
252 utility_score = $clamped,
253 utility_updates += 1,
254 updated_at = time::now()
255 "#,
256 )
257 .bind(("id", entity_id.to_string()))
258 .bind(("alpha", alpha))
259 .bind(("reward", reward))
260 .await?;
261
262 Ok(())
263}
264
265pub async fn get_utility_score(db: &Surreal<Db>, entity_id: &str) -> Result<f64, GraphError> {
267 #[derive(Deserialize)]
268 struct Row {
269 #[serde(default = "default_util")]
270 utility_score: f64,
271 }
272
273 fn default_util() -> f64 {
274 DEFAULT_UTILITY
275 }
276
277 let mut response = db
278 .query("SELECT utility_score FROM type::record($id)")
279 .bind(("id", entity_id.to_string()))
280 .await?;
281
282 let rows: Vec<Row> = super::deserialize_take(&mut response, 0)?;
283
284 Ok(rows
285 .first()
286 .map(|r| r.utility_score)
287 .unwrap_or(DEFAULT_UTILITY))
288}
289
290#[derive(Debug, Clone, Default)]
292pub struct ContributionStats {
293 pub total_contributions: u32,
294 pub successes: u32,
295 pub partials: u32,
296 pub failures: u32,
297 pub times_used: u32,
298 pub times_ignored: u32,
299}
300
301#[cfg(test)]
302mod tests {
303 use super::*;
304
305 #[test]
306 fn outcome_kind_reward_values() {
307 assert_eq!(OutcomeKind::Success.reward(), 1.0);
308 assert_eq!(OutcomeKind::Partial.reward(), 0.5);
309 assert_eq!(OutcomeKind::Failed.reward(), 0.0);
310 }
311
312 #[test]
313 fn outcome_kind_roundtrip() {
314 for kind in [
315 OutcomeKind::Success,
316 OutcomeKind::Partial,
317 OutcomeKind::Failed,
318 ] {
319 let s = kind.to_string();
320 let parsed: OutcomeKind = s.parse().unwrap();
321 assert_eq!(parsed, kind);
322 }
323 assert!("unknown".parse::<OutcomeKind>().is_err());
324 }
325
326 #[test]
327 fn ema_update_math() {
328 let current: f64 = 0.5;
329 let alpha: f64 = 0.1;
330
331 let success = (1.0 - alpha) * current + alpha * 1.0;
332 assert!((success - 0.55).abs() < 0.001);
333
334 let partial = (1.0 - alpha) * current + alpha * 0.5;
335 assert!((partial - 0.5).abs() < 0.001);
336
337 let failed = (1.0 - alpha) * current + alpha * 0.0;
338 assert!((failed - 0.45).abs() < 0.001);
339 }
340
341 #[test]
342 fn ema_converges() {
343 let mut score = 0.5;
344 for _ in 0..50 {
345 score = (1.0 - USED_ALPHA) * score + USED_ALPHA * 1.0;
346 }
347 assert!(score > 0.99);
348
349 let mut score = 0.5;
350 for _ in 0..50 {
351 score = (1.0 - USED_ALPHA) * score + USED_ALPHA * 0.0;
352 }
353 assert!(score < 0.01);
354 }
355
356 #[test]
357 fn unused_entity_gets_weaker_signal() {
358 let current = 0.5;
359 let used_step = (1.0 - USED_ALPHA) * current + USED_ALPHA * 1.0;
360 let unused_step = (1.0 - UNUSED_ALPHA) * current + UNUSED_ALPHA * UNUSED_REWARD;
361
362 assert!(used_step > current);
363 assert!(unused_step < current);
364 }
365}