""" Counterparty ecology — diverse adversarial agents for self-play. Each agent is a frozen policy with legal_actions() and rollout_action(). The ecology is designed so that a pure quote gets picked off by toxic takers, and only a mixed distribution survives. """ from __future__ import annotations import random from dataclasses import dataclass from typing import Optional, Protocol, Tuple from malkhut.state import ( ActionKind, AgentRole, FulfilmentPolicyParams, MarketWorldState, Side, ) from malkhut.actions import CounterpartyAction class CounterpartyPolicy(Protocol): role: AgentRole def legal_actions( self, state: MarketWorldState, params: Optional[FulfilmentPolicyParams] = None, ) -> Tuple[CounterpartyAction, ...]: ... def rollout_action( self, state: MarketWorldState, rng: random.Random, ) -> CounterpartyAction: ... @dataclass(frozen=True, slots=True) class ToxicTakerPolicy: role: AgentRole = AgentRole.TOXIC_TAKER sensitivity: float = 0.5 # LOWER = more aggressive (attacks more often) def legal_actions(self, state: MarketWorldState, params: Optional[FulfilmentPolicyParams] = None) -> Tuple[CounterpartyAction, ...]: return ( CounterpartyAction(self.role, ActionKind.NOOP, None, 0, 0.0), CounterpartyAction(self.role, ActionKind.CROSS_SPREAD, Side.BUY, 0, 0.25, toxicity=0.8), CounterpartyAction(self.role, ActionKind.CROSS_SPREAD, Side.SELL, 0, 0.25, toxicity=0.8), ) def rollout_action(self, state: MarketWorldState, rng: random.Random) -> CounterpartyAction: tox = state.trade_path.orderflow_toxicity if state.trade_path else 0.0 if tox > self.sensitivity or rng.random() < 0.3: # 30% base attack rate side = Side.SELL if (state.trade_path and state.trade_path.cross_venue_lead_score < 0) else Side.BUY return CounterpartyAction(self.role, ActionKind.CROSS_SPREAD, side, 0, 0.25, toxicity=tox) return CounterpartyAction(self.role, ActionKind.NOOP, None, 0, 0.0) @dataclass(frozen=True, slots=True) class PassiveMakerPolicy: role: AgentRole = AgentRole.PASSIVE_MAKER join_probability: float = 0.60 def legal_actions(self, state: MarketWorldState, params: Optional[FulfilmentPolicyParams] = None) -> Tuple[CounterpartyAction, ...]: return ( CounterpartyAction(self.role, ActionKind.NOOP, None, 0, 0.0), CounterpartyAction(self.role, ActionKind.PLACE, Side.BUY, 0, 0.20), CounterpartyAction(self.role, ActionKind.PLACE, Side.SELL, 0, 0.20), CounterpartyAction(self.role, ActionKind.CANCEL, None, 0, 0.0), ) def rollout_action(self, state: MarketWorldState, rng: random.Random) -> CounterpartyAction: if rng.random() < self.join_probability: side = Side.BUY if rng.random() < 0.5 else Side.SELL return CounterpartyAction(self.role, ActionKind.PLACE, side, 0, 0.20) return CounterpartyAction(self.role, ActionKind.NOOP, None, 0, 0.0) @dataclass(frozen=True, slots=True) class LatencyArbPolicy: role: AgentRole = AgentRole.LATENCY_ARB lead_threshold: float = 0.55 def legal_actions(self, state: MarketWorldState, params: Optional[FulfilmentPolicyParams] = None) -> Tuple[CounterpartyAction, ...]: return ( CounterpartyAction(self.role, ActionKind.NOOP, None, 0, 0.0), CounterpartyAction(self.role, ActionKind.CROSS_SPREAD, Side.BUY, 0, 0.15, toxicity=0.9), CounterpartyAction(self.role, ActionKind.CROSS_SPREAD, Side.SELL, 0, 0.15, toxicity=0.9), ) def rollout_action(self, state: MarketWorldState, rng: random.Random) -> CounterpartyAction: lead = state.trade_path.cross_venue_lead_score if state.trade_path else 0.0 if abs(lead) > self.lead_threshold: side = Side.BUY if lead > 0 else Side.SELL return CounterpartyAction(self.role, ActionKind.CROSS_SPREAD, side, 0, 0.15, toxicity=0.9) return CounterpartyAction(self.role, ActionKind.NOOP, None, 0, 0.0) @dataclass(frozen=True, slots=True) class NoiseTraderPolicy: role: AgentRole = AgentRole.NOISE_TRADER def legal_actions(self, state: MarketWorldState, params: Optional[FulfilmentPolicyParams] = None) -> Tuple[CounterpartyAction, ...]: return ( CounterpartyAction(self.role, ActionKind.NOOP, None, 0, 0.0), CounterpartyAction(self.role, ActionKind.CROSS_SPREAD, Side.BUY, 0, 0.05), CounterpartyAction(self.role, ActionKind.CROSS_SPREAD, Side.SELL, 0, 0.05), ) def rollout_action(self, state: MarketWorldState, rng: random.Random) -> CounterpartyAction: r = rng.random() if r < 0.10: return CounterpartyAction(self.role, ActionKind.CROSS_SPREAD, Side.BUY, 0, 0.05) if r < 0.20: return CounterpartyAction(self.role, ActionKind.CROSS_SPREAD, Side.SELL, 0, 0.05) return CounterpartyAction(self.role, ActionKind.NOOP, None, 0, 0.0) def default_counterparty_ecology() -> Tuple[CounterpartyPolicy, ...]: return ( ToxicTakerPolicy(), PassiveMakerPolicy(), LatencyArbPolicy(), NoiseTraderPolicy(), )