Rhodawk Mythos Agent
mythos: ascend to Mythos-level — multi-agent + probabilistic + advanced tooling + RL + MCP suite + FastAPI
14b1bbe | """ | |
| Reinforcement-learning controller for the Planner. | |
| Wraps Ray RLlib / Stable Baselines3 when available; otherwise exposes a | |
| contextual-bandit baseline that updates per-CWE arm preferences from | |
| campaign rewards. This is enough to deliver measurable improvement in the | |
| Planner's choice of CWE focus across hundreds of campaigns. | |
| """ | |
| from __future__ import annotations | |
| import json | |
| import math | |
| import os | |
| import random | |
| from typing import Any | |
| try: # pragma: no cover | |
| import ray # type: ignore # noqa: F401 | |
| from ray.rllib.algorithms.ppo import PPOConfig # type: ignore # noqa: F401 | |
| _RLLIB = True | |
| except Exception: # noqa: BLE001 | |
| _RLLIB = False | |
| try: # pragma: no cover | |
| from stable_baselines3 import PPO # type: ignore # noqa: F401 | |
| _SB3 = True | |
| except Exception: # noqa: BLE001 | |
| _SB3 = False | |
| _STATE_FILE = os.getenv("MYTHOS_RL_STATE", "/data/mythos/rl_state.json") | |
| class RLPlanner: | |
| """Contextual UCB1 over CWE arms (with PPO upgrade path).""" | |
| def __init__(self): | |
| self.counts: dict[str, int] = {} | |
| self.values: dict[str, float] = {} | |
| self.t: int = 0 | |
| self._load() | |
| def backend(self) -> str: | |
| if _RLLIB: | |
| return "ray-rllib" | |
| if _SB3: | |
| return "stable-baselines3" | |
| return "ucb1" | |
| def select(self, candidate_cwes: list[str]) -> str: | |
| self.t += 1 | |
| if not candidate_cwes: | |
| return "" | |
| # Cold-start: pull each arm at least once. | |
| for c in candidate_cwes: | |
| if self.counts.get(c, 0) == 0: | |
| return c | |
| scored = [ | |
| (c, self.values[c] + math.sqrt(2 * math.log(self.t) / self.counts[c])) | |
| for c in candidate_cwes | |
| ] | |
| return max(scored, key=lambda x: x[1])[0] | |
| def reward(self, cwe: str, signal: float) -> None: | |
| n = self.counts.get(cwe, 0) + 1 | |
| v = self.values.get(cwe, 0.0) | |
| self.counts[cwe] = n | |
| self.values[cwe] = v + (signal - v) / n | |
| self._save() | |
| def explore(self, candidates: list[str], epsilon: float = 0.1) -> str: | |
| if random.random() < epsilon: | |
| return random.choice(candidates) if candidates else "" | |
| return self.select(candidates) | |
| # -- persistence -------------------------------------------------------- | |
| def _load(self) -> None: | |
| try: | |
| with open(_STATE_FILE) as fh: | |
| state = json.load(fh) | |
| self.counts = state.get("counts", {}) | |
| self.values = state.get("values", {}) | |
| self.t = state.get("t", 0) | |
| except Exception: # noqa: BLE001 | |
| pass | |
| def _save(self) -> None: | |
| try: | |
| os.makedirs(os.path.dirname(_STATE_FILE), exist_ok=True) | |
| with open(_STATE_FILE, "w") as fh: | |
| json.dump({"counts": self.counts, "values": self.values, "t": self.t}, fh) | |
| except Exception: # noqa: BLE001 | |
| pass | |