From 90b824623ec86a5e5441af15292a60a9135f60d5 Mon Sep 17 00:00:00 2001 From: Frankie-Xu <92643488+Frankie-Xu@users.noreply.github.com> Date: Mon, 14 Sep 2026 03:54:35 +0800 Subject: [PATCH 1/2] fix: reject nonfinite Q updates before persistence --- memrl/service/value_driven.py | 25 ++++++++++++++++--------- 1 file changed, 16 insertions(+), 9 deletions(-) diff --git a/memrl/service/value_driven.py b/memrl/service/value_driven.py index 19b865d1f..510db6e02 100644 --- a/memrl/service/value_driven.py +++ b/memrl/service/value_driven.py @@ -17,6 +17,7 @@ from dataclasses import dataclass from typing import Any, Dict, List, Optional from datetime import datetime +from math import isfinite import random from memos.mem_os.main import MOS @@ -191,20 +192,26 @@ def update(self, memory_id: str, reward: float, next_max_q: Optional[float] = No old_meta = _meta_to_dict(getattr(item, "metadata", None)) old_q = float(old_meta.get("q_value", self.cfg.q_init_pos)) - target = float(reward) + (self.cfg.gamma * float(next_max_q or 0.0)) - new_q = (1.0 - self.cfg.alpha) * old_q + self.cfg.alpha * target + reward = float(reward) + next_q = 0.0 if next_max_q is None else float(next_max_q) + alpha, gamma = float(self.cfg.alpha), float(self.cfg.gamma) + old_ma = float(old_meta.get("reward_ma", 0.0)) + if not all(isfinite(value) for value in (old_q, reward, next_q, alpha, gamma, old_ma)): + raise ValueError("Q update inputs must be finite") + target = reward + gamma * next_q + new_q = (1.0 - alpha) * old_q + alpha * target + reward_ma = (1.0 - alpha) * old_ma + alpha * reward + if not all(isfinite(value) for value in (target, new_q, reward_ma)): + raise ValueError("Q update arithmetic must remain finite") # Optional Q floor: prevents Q from dropping below configured minimum. if getattr(self.cfg, "q_floor", None) is not None: - try: - new_q = max(float(self.cfg.q_floor), float(new_q)) - except Exception: - pass + floor = float(self.cfg.q_floor) + if not isfinite(floor): + raise ValueError("Q floor must be finite") + new_q = max(floor, new_q) visits = int(old_meta.get("q_visits", 0)) + 1 - # simple EMA for reward - old_ma = float(old_meta.get("reward_ma", 0.0)) - reward_ma = (1.0 - self.cfg.alpha) * old_ma + self.cfg.alpha * float(reward) new_meta = old_meta | { "q_value": float(new_q), From 3af564c26ad37ec7483ed7253e9129d3ddfa27fe Mon Sep 17 00:00:00 2001 From: Frankie-Xu <92643488+Frankie-Xu@users.noreply.github.com> Date: Mon, 14 Sep 2026 03:54:42 +0800 Subject: [PATCH 2/2] test: cover invalid Q inputs without external services --- tests/test_q_value_finite.py | 86 ++++++++++++++++++++++++++++++++++++ 1 file changed, 86 insertions(+) create mode 100644 tests/test_q_value_finite.py diff --git a/tests/test_q_value_finite.py b/tests/test_q_value_finite.py new file mode 100644 index 000000000..e626f0cca --- /dev/null +++ b/tests/test_q_value_finite.py @@ -0,0 +1,86 @@ +"""Exercise the real updater with an in-memory persistence boundary; no MemOS service.""" +import importlib.util +from pathlib import Path +import sys +from types import ModuleType, SimpleNamespace +import unittest +from unittest.mock import patch + + +def load_module(): + stub = ModuleType("memos.mem_os.main") + stub.MOS = object + spec = importlib.util.spec_from_file_location( + "value_driven_under_test", + Path(__file__).resolve().parents[1] / "memrl/service/value_driven.py", + ) + module = importlib.util.module_from_spec(spec) + with patch.dict(sys.modules, {"memos.mem_os.main": stub, spec.name: module}): + spec.loader.exec_module(module) + return module + + +vd = load_module() + + +class Store: + def __init__(self, metadata): + self.metadata = metadata.copy() + self.writes = [] + + def get(self, memory_id): + return SimpleNamespace(memory="synthetic", metadata=self.metadata) + + def update(self, memory_id, value): + self.writes.append(value) + self.metadata = value["metadata"] + + +class FiniteUpdateTests(unittest.TestCase): + def make_updater(self, metadata=None, **config): + store = Store(metadata or {"q_value": 0.0}) + updater = vd.QValueUpdater(None, "test", vd.RLConfig(**config)) + updater._get_text_mem = lambda: store + return updater, store + + def test_nonfinite_inputs_never_write(self): + for value in (float("nan"), float("inf"), -float("inf")): + for field in ("reward", "next_max_q", "alpha", "gamma", "q_floor", "q_value", "reward_ma"): + with self.subTest(value=value, field=field): + metadata = {"q_value": 0.0} + config, args = {}, {"reward": 1.0} + if field in ("q_value", "reward_ma"): + metadata[field] = value + elif field in ("reward", "next_max_q"): + args[field] = value + else: + config[field] = value + updater, store = self.make_updater(metadata, **config) + with self.assertRaises(ValueError): + updater.update("m", **args) + self.assertEqual(store.writes, []) + + def test_finite_overflow_never_writes(self): + cases = [ + ({"gamma": 2.0}, {"q_value": 0.0}, 1e308, 1e308), + ({"alpha": 2.0}, {"q_value": 0.0}, 1e308, None), + ({"alpha": 2.0}, {"q_value": 0.0, "reward_ma": -1e308}, 1e308, None), + ] + for config, metadata, reward, next_q in cases: + with self.subTest(config=config, metadata=metadata): + updater, store = self.make_updater(metadata, **config) + with self.assertRaises(ValueError): + updater.update("m", reward, next_q) + self.assertEqual(store.writes, []) + + def test_valid_repeated_updates_and_floor(self): + updater, store = self.make_updater(alpha=0.5, gamma=0.5, q_floor=-1.0) + self.assertEqual(updater.update("m", -2.0), -1.0) + self.assertEqual(updater.update("m", 2.0, 2.0), 1.0) + self.assertEqual(store.metadata["q_visits"], 2) + self.assertEqual(store.metadata["reward_ma"], 0.5) + self.assertEqual(store.metadata["last_reward"], 2.0) + + +if __name__ == "__main__": + unittest.main()