use serde_json::{json, Value};
use velesdb_core::agent::{FixedRate, ReinforcementContext, ReinforcementStrategy};
use super::{MemoryService, Metadata};
use crate::embedder::Embedder;
use crate::error::MemoryError;
use crate::storage::MemoryStore;
pub(crate) const RL_CONFIDENCE_KEY: &str = "_veles_rl_confidence";
const RL_SUCCESS_KEY: &str = "_veles_rl_success";
const RL_FAILURE_KEY: &str = "_veles_rl_failure";
pub(crate) const RL_NEUTRAL_CONFIDENCE: f32 = 0.5;
const RL_RERANK_WEIGHT: f32 = 0.5;
type Hit = (u64, f32, String);
type RankedHit = (Hit, Option<Metadata>, f32);
type RerankedHits = (Vec<Hit>, Vec<Option<Metadata>>);
impl<E: Embedder, S: MemoryStore> MemoryService<E, S> {
pub fn feedback(&self, id: u64, success: bool) -> Result<f32, MemoryError> {
let payload = self
.store
.get_metadata(id)?
.ok_or(MemoryError::UnknownMemory(id))?;
let confidence = read_confidence(&payload);
let mut success_count = read_count(&payload, RL_SUCCESS_KEY);
let mut failure_count = read_count(&payload, RL_FAILURE_KEY);
if success {
success_count += 1;
} else {
failure_count += 1;
}
let total = success_count + failure_count;
let mut context = ReinforcementContext::new().with_usage_count(total);
if let Some(rate) = success_rate(success_count, total) {
context = context.with_success_rate(rate);
}
let new_confidence = FixedRate::default().update_confidence(confidence, success, &context);
let mut updates = Metadata::new();
updates.insert(RL_CONFIDENCE_KEY.to_owned(), json!(new_confidence));
updates.insert(RL_SUCCESS_KEY.to_owned(), json!(success_count));
updates.insert(RL_FAILURE_KEY.to_owned(), json!(failure_count));
self.store.update_metadata(id, &updates)?;
Ok(new_confidence)
}
pub(crate) fn rl_rerank(hits: Vec<Hit>, payloads: Vec<Option<Metadata>>) -> RerankedHits {
if hits.len() < 2 {
return (hits, payloads);
}
let mut ranked: Vec<RankedHit> = hits
.into_iter()
.zip(payloads)
.map(|(hit, payload)| {
let confidence = payload
.as_ref()
.map_or(RL_NEUTRAL_CONFIDENCE, read_confidence);
let blended = blended_score(hit.1, confidence);
(hit, payload, blended)
})
.collect();
ranked.sort_by(|a, b| b.2.total_cmp(&a.2));
let mut out_hits = Vec::with_capacity(ranked.len());
let mut out_payloads = Vec::with_capacity(ranked.len());
for (hit, payload, _) in ranked {
out_hits.push(hit);
out_payloads.push(payload);
}
(out_hits, out_payloads)
}
}
fn blended_score(similarity: f32, confidence: f32) -> f32 {
let base = f32::midpoint(similarity, 1.0);
let factor = 1.0 + RL_RERANK_WEIGHT * (2.0 * confidence - 1.0);
base * factor
}
#[allow(
clippy::cast_possible_truncation,
reason = "confidence is a bounded [0,1] weight; f64→f32 rounding is immaterial and the result is clamped"
)]
fn read_confidence(payload: &Metadata) -> f32 {
payload
.get(RL_CONFIDENCE_KEY)
.and_then(Value::as_f64)
.map_or(RL_NEUTRAL_CONFIDENCE, |v| (v as f32).clamp(0.0, 1.0))
}
fn read_count(payload: &Metadata, key: &str) -> u64 {
payload.get(key).and_then(Value::as_u64).unwrap_or(0)
}
#[allow(
clippy::cast_precision_loss,
reason = "feedback tallies are small counters; an approximate rate is all the strategy needs"
)]
fn success_rate(success_count: u64, total: u64) -> Option<f32> {
if total == 0 {
None
} else {
Some(success_count as f32 / total as f32)
}
}
#[cfg(all(test, feature = "persistence"))]
mod tests {
use crate::embedder::HashEmbedder;
use crate::service::MemoryService;
use crate::DEFAULT_DIMENSION;
use tempfile::TempDir;
fn service() -> (TempDir, MemoryService<HashEmbedder>) {
let dir = TempDir::new().expect("tempdir");
let embedder = HashEmbedder::new(DEFAULT_DIMENSION);
let svc = MemoryService::open(dir.path(), embedder).expect("open store");
(dir, svc)
}
#[test]
fn feedback_raises_confidence_on_success_and_lowers_on_failure() {
let (_dir, svc) = service();
let id = svc.remember("rust prevents data races", &[], None).unwrap();
let up = svc.feedback(id, true).unwrap();
assert!(up > 0.5, "success should raise confidence, got {up}");
let down = svc.feedback(id, false).unwrap();
assert!(down < up, "failure should lower confidence, got {down}");
}
#[test]
fn feedback_is_clamped_and_monotonic_under_repeated_success() {
let (_dir, svc) = service();
let id = svc.remember("clamp me", &[], None).unwrap();
let mut last = 0.5_f32;
for _ in 0..50 {
let c = svc.feedback(id, true).unwrap();
assert!(c >= last - f32::EPSILON, "confidence must not decrease");
assert!(c <= 1.0, "confidence must stay clamped to 1.0, got {c}");
last = c;
}
assert!(
last > 0.99,
"many successes should saturate near 1.0, got {last}"
);
}
#[test]
fn feedback_persists_across_reopen() {
let dir = TempDir::new().expect("tempdir");
let id;
let after;
{
let svc =
MemoryService::open(dir.path(), HashEmbedder::new(DEFAULT_DIMENSION)).unwrap();
id = svc.remember("durable confidence", &[], None).unwrap();
svc.feedback(id, true).unwrap();
after = svc.feedback(id, true).unwrap();
}
let svc = MemoryService::open(dir.path(), HashEmbedder::new(DEFAULT_DIMENSION)).unwrap();
let resumed = svc.feedback(id, true).unwrap();
assert!(
resumed > after,
"confidence must resume from persisted {after}, got {resumed}"
);
}
#[test]
fn feedback_teaches_recall_to_prefer_the_authoritative_answer() {
let (_dir, svc) = service();
svc.remember(
"Use `Client::builder().timeout(d).build()` to configure the HTTP client timeout",
&[],
None,
)
.unwrap();
svc.remember(
"Deprecated: set the HTTP client timeout via the global `CLIENT_TIMEOUT` env var",
&[],
None,
)
.unwrap();
let query = "how to configure the http client timeout";
let baseline = svc.recall(query, 2, None).unwrap();
assert_eq!(baseline.len(), 2, "both facts should be recalled");
let authoritative = baseline[1].id; let deprecated = baseline[0].id;
for _ in 0..15 {
svc.feedback(authoritative, true).unwrap();
svc.feedback(deprecated, false).unwrap();
}
let after = svc.recall(query, 2, None).unwrap();
assert_eq!(
after[0].id, authoritative,
"recall must now lead with the fact the team kept marking useful"
);
let sim_before = baseline
.iter()
.find(|r| r.id == authoritative)
.unwrap()
.score;
let sim_after = after.iter().find(|r| r.id == authoritative).unwrap().score;
assert!(
(sim_before - sim_after).abs() < 1e-6,
"feedback re-orders results; it must not fabricate a different similarity score"
);
}
#[test]
fn recall_order_is_untouched_without_feedback() {
let (_dir, svc) = service();
for fact in ["alpha fact", "beta fact", "gamma fact", "delta fact"] {
svc.remember(fact, &[], None).unwrap();
}
let a = svc.recall("fact", 4, None).unwrap();
let b = svc.recall("fact", 4, None).unwrap();
let ids_a: Vec<u64> = a.iter().map(|r| r.id).collect();
let ids_b: Vec<u64> = b.iter().map(|r| r.id).collect();
assert_eq!(ids_a, ids_b, "recall must be deterministic and unreordered");
}
#[test]
fn feedback_on_unknown_id_errors() {
let (_dir, svc) = service();
assert!(svc.feedback(999, true).is_err(), "unknown id must error");
}
#[test]
fn blend_never_inverts_ranking_even_on_negative_similarity() {
use super::blended_score;
for &sim in &[-0.99_f32, -0.5, -0.12, 0.0, 0.3, 0.95] {
let punished = blended_score(sim, 0.0);
let neutral = blended_score(sim, 0.5);
let reinforced = blended_score(sim, 1.0);
assert!(
reinforced >= neutral && neutral >= punished,
"sim={sim}: confidence inverted the ranking ({punished} <= {neutral} <= {reinforced})"
);
}
assert!(
blended_score(-0.12, 1.0) > blended_score(-0.10, 0.5),
"reinforcement must overcome a small similarity gap even when negative"
);
}
}