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::FactStore;
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: FactStore> MemoryService<E, S> {
pub fn feedback(&self, id: u64, success: bool) -> Result<f32, MemoryError> {
let _generation = self.enter_generation();
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"
)]
pub(crate) 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"))]
#[path = "reinforce_tests.rs"]
mod tests;