Skip to main content

recall_echo/graph/
utility.rs

1//! Outcome feedback loop for adaptive entity learning.
2//!
3//! Tracks which graph entities contributed to session outcomes (success/partial/failure)
4//! and adjusts their `utility_score` via exponential moving average.
5//!
6//! Phase 1 of Adaptive Entity Learning v2.
7
8use serde::{Deserialize, Serialize};
9use surrealdb::Surreal;
10
11use super::error::GraphError;
12use super::store::Db;
13
14/// The result of a task or session outcome.
15#[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    /// Numeric reward signal for EMA update.
25    #[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
58/// Default utility score for new entities.
59pub const DEFAULT_UTILITY: f64 = 0.5;
60
61/// EMA alpha for entities that were retrieved AND used.
62const USED_ALPHA: f64 = 0.1;
63
64/// Smaller EMA alpha for entities that were retrieved but not used.
65const UNUSED_ALPHA: f64 = 0.05;
66
67/// Reward override for "retrieved but not used" — slight negative signal.
68const UNUSED_REWARD: f64 = 0.3;
69
70/// Report from a feedback recording operation.
71#[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
79/// Record outcome feedback: link retrieved entities to an outcome and update utility scores.
80pub 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    // Build a HashSet for O(1) "was used" lookups instead of O(n) per entity
99    let used_set: Option<std::collections::HashSet<&str>> =
100        used_entity_ids.map(|ids| ids.iter().map(|s| s.as_str()).collect());
101
102    // Process all entities concurrently — each entity's feedback is independent
103    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
238/// Atomic EMA update — single query, no read-modify-write race.
239async fn update_utility_score(
240    db: &Surreal<Db>,
241    entity_id: &str,
242    reward: f64,
243    alpha: f64,
244) -> Result<(), GraphError> {
245    // Inline EMA: new = (1 - alpha) * current + alpha * reward, clamped to [0, 1].
246    // SurrealDB doesn't have math::clamp, so we use nested IF expressions.
247    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
265/// Get the current utility score for an entity.
266pub 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/// Get aggregate contribution stats for an entity.
291#[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}