diff --git a/.gitignore b/.gitignore index 6a16f7d5..df2ad42e 100644 --- a/.gitignore +++ b/.gitignore @@ -190,7 +190,10 @@ conf/ # IDE files .vscode/ .idea/ -improvements +# Local backlog of the /improvements and /design-feature workflows: working +# notes for the burndown runs, not something the repo should carry. +improvements/ +features/ bots/gateway-files/ .DS_Store Dockerfile.patched-hummingbot diff --git a/README.md b/README.md index fe429565..df1a5882 100644 --- a/README.md +++ b/README.md @@ -170,9 +170,15 @@ The `.env` file contains all configuration. Key settings: USERNAME=admin # API username PASSWORD=admin # API password CONFIG_PASSWORD=admin # Encrypts bot credentials +DEBUG_MODE=false # Verbose logging and reload DATABASE_URL=... # PostgreSQL connection GATEWAY_URL=... # Gateway URL (for DEX) +# Performance snapshots and backtests +PERFORMANCE_EXECUTOR_SNAPSHOT_INTERVAL=60 # Seconds between live executor snapshots +PERFORMANCE_RETENTION_DAYS=0 # Delete snapshots older than N days; 0 keeps them forever +BACKTESTING_MAX_CONCURRENT=1 # Backtests allowed to run at once (one core each) + # Tailscale (recommended for production) TAILSCALE_ENABLED=true TAILSCALE_AUTH_KEY=tskey-auth-... @@ -186,6 +192,11 @@ TAILSCALE_HOSTNAME=hummingbot-api # MagicDNS hostname on your tailnet # DB_BIND=127.0.0.1 ``` +These are the settings most deployments touch, not the full list: `config.py` is the authoritative +list of every setting, its default and what it does. The `.env` that `setup.sh` generates also carries +the optional `PERFORMANCE_`, `BACKTESTING_`, `MARKET_DATA_`, `CORS_` and `AWS_` groups as commented-out +lines showing their defaults, so you can see and override them without leaving the file. + Edit `.env` and restart with `make deploy` to apply changes. ## Secure Connection via Tailscale diff --git a/bots/controllers/generic/lp_rebalancer/lp_rebalancer.py b/bots/controllers/generic/lp_rebalancer/lp_rebalancer.py index 7a774c03..ba3104d1 100644 --- a/bots/controllers/generic/lp_rebalancer/lp_rebalancer.py +++ b/bots/controllers/generic/lp_rebalancer/lp_rebalancer.py @@ -93,6 +93,19 @@ class LPRebalancerConfig(ControllerConfigBase): description="Extra % to swap beyond deficit to account for slippage (e.g., 0.01 = 0.01%)" ) + @field_validator("lp_provider") + @classmethod + def validate_lp_provider(cls, v: str) -> str: + """The trading type is never guessed: Gateway rejects a guessed one with a 400, so an + untyped provider has to fail at config load rather than mid-operation. The core's + parse_provider defaults an untyped provider to "router", which is the wrong branch + entirely for an LP controller, so the contract is enforced here instead.""" + if "/" not in v: + raise ValueError( + f"Invalid lp_provider '{v}': expected 'name/type' (e.g. 'meteora/clmm')" + ) + return v + @field_validator("sell_price_min", "sell_price_max", "buy_price_min", "buy_price_max", mode="before") @classmethod def validate_price_limits(cls, v): @@ -173,7 +186,9 @@ def __init__(self, config: LPRebalancerConfig, *args, **kwargs): super().__init__(config, *args, **kwargs) self.config: LPRebalancerConfig = config - # Parse lp_provider into dex_name and trading_type for gateway calls + # Parse lp_provider into dex_name and trading_type for gateway calls. The config + # validator guarantees the "name/type" form, so parse_provider's own default for an + # untyped provider is never reached. self.lp_dex_name, self.lp_trading_type = parse_provider(config.lp_provider) # Parse token symbols from trading pair diff --git a/config.py b/config.py index 042848ba..f667b44b 100644 --- a/config.py +++ b/config.py @@ -55,6 +55,22 @@ class MarketDataSettings(BaseSettings): default=60.0, description="Maximum allowed WebSocket subscription update interval in seconds" ) + ws_executor_min_update_interval: float = Field( + default=0.5, + description="Minimum allowed /ws/executors subscription update interval in seconds. " + "The floor is stricter than the market-data one because executor push loops " + "hit the database (executors, performance reports, positions with per-position " + "rate lookups) instead of reading in-memory candles and order books" + ) + ws_executor_max_update_interval: float = Field( + default=60.0, + description="Maximum allowed /ws/executors subscription update interval in seconds" + ) + ws_executor_default_update_interval: float = Field( + default=2.0, + description="Update interval applied to a /ws/executors subscription that does not " + "request one, in seconds" + ) ticker_update_interval: int = Field( default=30, description="How often to refresh tickers from connected exchanges in seconds" @@ -211,13 +227,28 @@ class AppSettings(BaseSettings): class BacktestingSettings(BaseSettings): - """Backtest result retention. + """Backtest execution limits and result retention. A finished backtest is ~98% bulk arrays (processed_data, pnl_timeseries) and only a few KB of metrics, so full payloads are archived to disk and only metrics stay resident. Retention is therefore a count of results, not a memory budget. + + A run executes in its own worker process and saturates a core for its whole duration, + so the two execution limits are about the box, not about memory: how many cores runs + may claim at once, and how long one is allowed to claim one before being abandoned. """ + max_concurrent: int = Field( + default=1, + description=( + "How many backtests may run at once; further submissions queue. Runs are isolated " + "in separate processes, so this can be raised up to the cores you are willing to give them" + ) + ) + timeout_seconds: float = Field( + default=1800.0, + description="Wall-clock budget for one backtest; the worker is killed and the task fails past it" + ) max_results: int = Field( default=100, description="How many finished backtests to retain before the oldest are reaped" @@ -226,10 +257,48 @@ class BacktestingSettings(BaseSettings): default="bots/data/backtests", description="Directory holding archived backtest payloads (inside the bots volume, so it survives redeploys)" ) + candles_cache_path: str = Field( + default="bots/data/backtests/candles", + description="Directory holding downloaded candle history shared by backtest workers" + ) + candles_cache_entries: int = Field( + default=32, + description=( + "How many downloaded candle ranges to keep; the least recently used are dropped past it. " + "0 disables the cache and makes every run download its own history again" + ) + ) + candles_cache_ttl_seconds: float = Field( + default=3600.0, + description=( + "How long a downloaded candle range may be reused. A window ending near now is fetched " + "with its last candle still forming, so an entry is refetched once it is older than this" + ) + ) model_config = SettingsConfigDict(env_prefix="BACKTESTING_", extra="ignore") +class PerformanceSettings(BaseSettings): + """Performance snapshot cadence and retention.""" + + executor_snapshot_interval: int = Field( + default=60, + description="How often a live executor's performance is snapshotted, in seconds. " + "Finer than the controller dump because executors are short-lived: at " + "a 5-minute grain a three-minute position executor gets one point." + ) + retention_days: int = Field( + default=0, + description="Delete performance snapshots (executor AND controller) older than " + "this many days. 0 keeps everything forever, which is what every " + "existing deployment does today -- an upgrade must not start deleting " + "an operator's history." + ) + + model_config = SettingsConfigDict(env_prefix="PERFORMANCE_", extra="ignore") + + class Settings(BaseSettings): """Combined application settings.""" @@ -242,6 +311,7 @@ class Settings(BaseSettings): cors: CORSSettings = Field(default_factory=CORSSettings) app: AppSettings = Field(default_factory=AppSettings) backtesting: BacktestingSettings = Field(default_factory=BacktestingSettings) + performance: PerformanceSettings = Field(default_factory=PerformanceSettings) # Direct banned_tokens field to handle env parsing banned_tokens: List[str] = Field( diff --git a/database/__init__.py b/database/__init__.py index 7f759fb4..637a8390 100644 --- a/database/__init__.py +++ b/database/__init__.py @@ -4,6 +4,7 @@ Base, BotRun, ControllerPerformanceSnapshot, + ExecutorPerformanceSnapshot, FundingPayment, GatewayCLMMEvent, GatewayCLMMPosition, @@ -17,6 +18,7 @@ AccountRepository, BotRunRepository, ControllerPerformanceRepository, + ExecutorPerformanceRepository, ExecutorRepository, FundingRepository, GatewayCLMMRepository, @@ -28,10 +30,10 @@ __all__ = [ "AccountState", "TokenState", "Order", "Trade", "PositionSnapshot", "FundingPayment", "BotRun", "GatewaySwap", "GatewayCLMMPosition", "GatewayCLMMEvent", - "ControllerPerformanceSnapshot", + "ControllerPerformanceSnapshot", "ExecutorPerformanceSnapshot", "Base", "AsyncDatabaseManager", "AccountRepository", "BotRunRepository", "ControllerPerformanceRepository", - "ExecutorRepository", + "ExecutorPerformanceRepository", "ExecutorRepository", "OrderRepository", "TradeRepository", "FundingRepository", "GatewaySwapRepository", "GatewayCLMMRepository" ] diff --git a/database/models.py b/database/models.py index 18bf9322..3c74efe1 100644 --- a/database/models.py +++ b/database/models.py @@ -1,4 +1,4 @@ -from sqlalchemy import TIMESTAMP, Column, ForeignKey, Integer, Numeric, String, Text, UniqueConstraint, func +from sqlalchemy import TIMESTAMP, Boolean, Column, ForeignKey, Index, Integer, Numeric, String, Text, UniqueConstraint, func from sqlalchemy.ext.declarative import declarative_base from sqlalchemy.orm import relationship @@ -453,6 +453,52 @@ class ControllerPerformanceSnapshot(Base): custom_info = Column(Text, nullable=True) # JSON dict of custom info +class ExecutorPerformanceSnapshot(Base): + """Periodic snapshot of a live executor's performance, plus one terminal row. + + Written by ExecutorService: on the snapshot tick for everything in + _active_executors, and once more at completion with is_terminal=True. The terminal + row is what makes a closed executor's series answerable from this table alone -- + no join to ExecutorRecord, and no exposure to the two paths that leave that record's + metrics at their creation-time zeros (the startup reap, and the except branch in + _persist_executor_completed). + + Deliberately narrow and typed, unlike controller_performance_snapshots: ExecutorInfo + is named field by field all over this repo, while the controller's PerformanceReport + is versioned by the core and has to be blobbed. Typing it also keeps the heavy + custom_info payloads (fill_events, levels_by_state, ...) out of a per-minute row. + """ + __tablename__ = "executor_performance_snapshots" + __table_args__ = ( + # The hot query is WHERE executor_id = ? ORDER BY timestamp DESC -- one executor's + # series, which is what both the history route and the reap lookup ask for. + Index("ix_exec_perf_executor_timestamp", "executor_id", "timestamp"), + ) + + id = Column(Integer, primary_key=True, index=True) + timestamp = Column(TIMESTAMP(timezone=True), server_default=func.now(), nullable=False, index=True) + + # Identity, denormalized from ExecutorRecord so a series needs no join. + executor_id = Column(String, nullable=False, index=True) + executor_type = Column(String, nullable=False, index=True) + account_name = Column(String, nullable=False, index=True) + connector_name = Column(String, nullable=False) + trading_pair = Column(String, nullable=False) + controller_id = Column(String, nullable=False, default="main", index=True) + + status = Column(String, nullable=False) # RunnableStatus name + close_type = Column(String, nullable=True) # only ever set on the terminal row + is_terminal = Column(Boolean, nullable=False, default=False, index=True) + + # The four ExecutorInfo metrics, same precision as ExecutorRecord. There is NO + # separate volume column: filled_amount_quote IS the volume traded, on every executor + # type including LP -- see test_executor_volume_is_the_filled_amount.py. + net_pnl_quote = Column(Numeric(precision=30, scale=18), nullable=False, default=0) + net_pnl_pct = Column(Numeric(precision=10, scale=6), nullable=False, default=0) + cum_fees_quote = Column(Numeric(precision=30, scale=18), nullable=False, default=0) + filled_amount_quote = Column(Numeric(precision=30, scale=18), nullable=False, default=0) + + class ExecutorRecord(Base): """Database model for executor state persistence.""" __tablename__ = "executors" diff --git a/database/repositories/__init__.py b/database/repositories/__init__.py index 11d19f88..d932c649 100644 --- a/database/repositories/__init__.py +++ b/database/repositories/__init__.py @@ -1,6 +1,7 @@ from .account_repository import AccountRepository from .bot_run_repository import BotRunRepository from .controller_performance_repository import ControllerPerformanceRepository +from .executor_performance_repository import ExecutorPerformanceRepository from .executor_repository import ExecutorRepository from .funding_repository import FundingRepository from .gateway_amm_repository import GatewayAMMRepository @@ -13,6 +14,7 @@ "AccountRepository", "BotRunRepository", "ControllerPerformanceRepository", + "ExecutorPerformanceRepository", "ExecutorRepository", "FundingRepository", "OrderRepository", diff --git a/database/repositories/bot_run_repository.py b/database/repositories/bot_run_repository.py index f6566cc4..f9edc67a 100644 --- a/database/repositories/bot_run_repository.py +++ b/database/repositories/bot_run_repository.py @@ -36,13 +36,12 @@ async def create_bot_run( deployment_status="DEPLOYED", run_status="CREATED" ) - + self.session.add(bot_run) await self.session.flush() await self.session.refresh(bot_run) return bot_run - async def update_bot_run_stopped( self, bot_name: str, @@ -56,18 +55,24 @@ async def update_bot_run_stopped( or_(BotRun.run_status == "RUNNING", BotRun.run_status == "CREATED") ) ).order_by(desc(BotRun.deployed_at)) - + result = await self.session.execute(stmt) bot_run = result.scalar_one_or_none() - + if bot_run: bot_run.run_status = "STOPPED" if not error_message else "ERROR" - bot_run.stopped_at = datetime.utcnow() + # Aware UTC, not utcnow(): stopped_at is TIMESTAMP(timezone=True), and a naive + # datetime is stored as if it were already in the session's local timezone, so + # utcnow() landed the row at the server's UTC offset behind the real stop time. + # stop-and-archive masked it -- update_bot_run_archived overwrote the value + # with a correct one -- but a bot stopped and never archived kept the skew, + # and run duration and performance-window attribution are read off this field. + bot_run.stopped_at = datetime.now(timezone.utc) bot_run.final_status = json.dumps(final_status) if final_status else None bot_run.error_message = error_message await self.session.flush() await self.session.refresh(bot_run) - + return bot_run async def update_bot_run_archived(self, bot_name: str) -> Optional[BotRun]: @@ -75,16 +80,16 @@ async def update_bot_run_archived(self, bot_name: str) -> Optional[BotRun]: stmt = select(BotRun).where( BotRun.bot_name == bot_name ).order_by(desc(BotRun.deployed_at)) - + result = await self.session.execute(stmt) bot_run = result.scalar_one_or_none() - + if bot_run: bot_run.deployment_status = "ARCHIVED" bot_run.stopped_at = datetime.now(timezone.utc) await self.session.flush() await self.session.refresh(bot_run) - + return bot_run async def get_bot_runs( @@ -100,7 +105,7 @@ async def get_bot_runs( ) -> List[BotRun]: """Get bot runs with optional filters.""" stmt = select(BotRun) - + conditions = [] if bot_name: conditions.append(BotRun.bot_name == bot_name) @@ -114,12 +119,12 @@ async def get_bot_runs( conditions.append(BotRun.run_status == run_status) if deployment_status: conditions.append(BotRun.deployment_status == deployment_status) - + if conditions: stmt = stmt.where(and_(*conditions)) - + stmt = stmt.order_by(desc(BotRun.deployed_at)).limit(limit).offset(offset) - + result = await self.session.execute(stmt) return result.scalars().all() @@ -134,7 +139,7 @@ async def get_latest_bot_run(self, bot_name: str) -> Optional[BotRun]: stmt = select(BotRun).where( BotRun.bot_name == bot_name ).order_by(desc(BotRun.deployed_at)) - + result = await self.session.execute(stmt) return result.scalar_one_or_none() @@ -146,7 +151,7 @@ async def get_active_bot_runs(self) -> List[BotRun]: BotRun.deployment_status == "DEPLOYED" ) ).order_by(desc(BotRun.deployed_at)) - + result = await self.session.execute(stmt) return result.scalars().all() @@ -156,7 +161,7 @@ async def get_bot_run_stats(self) -> Dict[str, Any]: total_stmt = select(func.count(BotRun.id)) total_result = await self.session.execute(total_stmt) total_runs = total_result.scalar() - + # Active runs active_stmt = select(func.count(BotRun.id)).where( and_( @@ -166,7 +171,7 @@ async def get_bot_run_stats(self) -> Dict[str, Any]: ) active_result = await self.session.execute(active_stmt) active_runs = active_result.scalar() - + # Runs by strategy type strategy_stmt = select( BotRun.strategy_type, @@ -174,7 +179,7 @@ async def get_bot_run_stats(self) -> Dict[str, Any]: ).group_by(BotRun.strategy_type) strategy_result = await self.session.execute(strategy_stmt) strategy_counts = {row.strategy_type: row.count for row in strategy_result} - + # Runs by status status_stmt = select( BotRun.run_status, @@ -182,7 +187,7 @@ async def get_bot_run_stats(self) -> Dict[str, Any]: ).group_by(BotRun.run_status) status_result = await self.session.execute(status_stmt) status_counts = {row.run_status: row.count for row in status_result} - + return { "total_runs": total_runs, "active_runs": active_runs, @@ -215,4 +220,4 @@ async def delete_bot_runs_by_bot_name(self, bot_name: str) -> int: if count > 0: await self.session.flush() - return count \ No newline at end of file + return count diff --git a/database/repositories/controller_performance_repository.py b/database/repositories/controller_performance_repository.py index e7494c52..8b88c067 100644 --- a/database/repositories/controller_performance_repository.py +++ b/database/repositories/controller_performance_repository.py @@ -22,22 +22,27 @@ def _interval_to_minutes(interval: str) -> int: @staticmethod def _sample_by_interval(history: List[Dict], interval_minutes: int) -> List[Dict]: + """Thin a descending-timestamp series to one row per interval PER CONTROLLER. + + The cursor is kept per (bot_name, controller_id), not once for the whole result. + A single global cursor makes `interval` a rate limit on the merged series, so an + unnarrowed query over a fleet drops whole controllers instead of thinning each of + them -- a 12-controller fleet answered with 11 of 12 at `1h`. Input order is + preserved, so the result stays descending by timestamp across controllers. + """ if not history or interval_minutes <= 5: return history sampled = [] - last_sampled_time = None + last_sampled_time: Dict[Tuple[Optional[str], Optional[str]], datetime] = {} for item in history: + scope = (item.get("bot_name"), item.get("controller_id")) item_time = datetime.fromisoformat(item["timestamp"].replace('Z', '+00:00')) - if last_sampled_time is None: + previous = last_sampled_time.get(scope) + if previous is None or (previous - item_time).total_seconds() / 60 >= interval_minutes: sampled.append(item) - last_sampled_time = item_time - else: - time_diff = (last_sampled_time - item_time).total_seconds() / 60 - if time_diff >= interval_minutes: - sampled.append(item) - last_sampled_time = item_time + last_sampled_time[scope] = item_time return sampled diff --git a/database/repositories/executor_performance_repository.py b/database/repositories/executor_performance_repository.py new file mode 100644 index 00000000..9f5e3803 --- /dev/null +++ b/database/repositories/executor_performance_repository.py @@ -0,0 +1,320 @@ +"""Repository for the executor performance snapshot series. + +Modelled on ControllerPerformanceRepository -- same descending-timestamp ordering, same +cursor semantics, same over-fetch-by-one has_more -- so a client pages either subject of +/performance/history with identical code. + +The one structural difference is the grain. The controller repository hard-codes a +5-minute grain in two places because that is what its dump loop writes; executors are +snapshotted far more often (60s by default, because a position executor can live three +minutes), so the sampler takes the grain as a parameter. The controller repository is +deliberately left alone rather than generalized: it sits on the wire-compatible +/bot-orchestration/controller-performance-* path. +""" +from datetime import datetime +from typing import Dict, List, Optional, Tuple + +from sqlalchemy import delete, desc, func, select +from sqlalchemy.ext.asyncio import AsyncSession + +from database.models import ControllerPerformanceSnapshot, ExecutorPerformanceSnapshot + +# How often a live executor is snapshotted, in minutes. Only used when the caller does +# not say; ExecutorService passes the configured interval through. +DEFAULT_GRAIN_MINUTES = 1.0 + + +class ExecutorPerformanceRepository: + def __init__(self, session: AsyncSession, grain_minutes: float = DEFAULT_GRAIN_MINUTES): + self.session = session + # Guard against a zero or negative interval reaching the sampler's divisor. + self.grain_minutes = grain_minutes if grain_minutes > 0 else DEFAULT_GRAIN_MINUTES + + @staticmethod + def _interval_to_minutes(interval: str) -> int: + interval_map = { + "1m": 1, "5m": 5, "15m": 15, "30m": 30, + "1h": 60, "4h": 240, "12h": 720, "1d": 1440 + } + return interval_map.get(interval, 5) + + @staticmethod + def _sample_by_interval(history: List[Dict], interval_minutes: int, grain_minutes: float) -> List[Dict]: + """Thin a descending-timestamp series down to one row per interval PER EXECUTOR. + + The cursor is kept per executor_id, not once for the whole result. A single global + cursor turns `interval` into a rate limit on the *merged* series: an unnarrowed + query interleaves every live executor's rows on the same grain, so the executor + that happens to own the newest row in each window survives and the rest are not + thinned but dropped entirely -- absent from a 200 response, indistinguishable from + executors that never reported. At the 60s write grain against the 5m default that + is roughly one executor kept per five snapshot rows, fleet-wide. + + Grouping makes the interval mean what the parameter says: each executor's own + series is thinned to its own resolution, and no scope disappears. Input order is + preserved, so the result stays descending by timestamp across executors. + + A no-op when the requested interval is no coarser than what is stored: there is + nothing to thin, and the caller gets the native grain. + + Thinning restarts at a page boundary -- the per-executor cursors do not survive in + `next_cursor`, which is a timestamp -- so an executor's first row on page two can + sit closer than `interval` to its last row on page one. That over-samples a scope + by at most one row per page; it never drops one. + """ + if not history or interval_minutes <= grain_minutes: + return history + + sampled = [] + last_sampled_time: Dict[Optional[str], datetime] = {} + + for item in history: + scope = item.get("executor_id") + item_time = datetime.fromisoformat(item["timestamp"].replace('Z', '+00:00')) + previous = last_sampled_time.get(scope) + if previous is None or (previous - item_time).total_seconds() / 60 >= interval_minutes: + sampled.append(item) + last_sampled_time[scope] = item_time + + return sampled + + async def save_snapshots(self, snapshots: List[Dict]) -> List[ExecutorPerformanceSnapshot]: + """Save a batch of executor performance snapshots with a single add_all/flush. + + Each item carries the identity columns, the status pair and the four metrics; + `snapshot_timestamp` is optional and defaults to the server clock. + """ + if not snapshots: + return [] + + rows = [] + for item in snapshots: + data = { + "executor_id": item["executor_id"], + "executor_type": item["executor_type"], + "account_name": item["account_name"], + "connector_name": item["connector_name"], + "trading_pair": item["trading_pair"], + "controller_id": item.get("controller_id") or "main", + "status": item["status"], + "close_type": item.get("close_type"), + "is_terminal": bool(item.get("is_terminal", False)), + "net_pnl_quote": item.get("net_pnl_quote", 0), + "net_pnl_pct": item.get("net_pnl_pct", 0), + "cum_fees_quote": item.get("cum_fees_quote", 0), + "filled_amount_quote": item.get("filled_amount_quote", 0), + } + if item.get("snapshot_timestamp"): + data["timestamp"] = item["snapshot_timestamp"] + rows.append(ExecutorPerformanceSnapshot(**data)) + + self.session.add_all(rows) + await self.session.flush() + return rows + + async def get_latest_for(self, executor_ids: List[str]) -> Dict[str, Dict]: + """The most recent snapshot of each of the given executors, keyed by executor_id. + + This is what lets the startup reap terminate an orphaned executor at its last + observed figures instead of the creation-time zeros. Executors with no snapshot + are simply absent from the result. + """ + if not executor_ids: + return {} + + latest = ( + select( + ExecutorPerformanceSnapshot.executor_id, + func.max(ExecutorPerformanceSnapshot.timestamp).label("max_timestamp"), + ) + .where(ExecutorPerformanceSnapshot.executor_id.in_(executor_ids)) + .group_by(ExecutorPerformanceSnapshot.executor_id) + .subquery() + ) + + query = ( + select(ExecutorPerformanceSnapshot) + .join( + latest, + (ExecutorPerformanceSnapshot.executor_id == latest.c.executor_id) & + (ExecutorPerformanceSnapshot.timestamp == latest.c.max_timestamp) + ) + ) + + result = await self.session.execute(query) + return {s.executor_id: self._to_dict(s) for s in result.scalars().all()} + + async def get_latest( + self, + executor_id: Optional[str] = None, + executor_type: Optional[str] = None, + controller_id: Optional[str] = None, + account_name: Optional[str] = None, + connector_name: Optional[str] = None, + trading_pair: Optional[str] = None, + limit: Optional[int] = None, + ) -> List[Dict]: + """The most recent snapshot of every matching executor, newest first. + + The executor counterpart of ControllerPerformanceRepository.get_latest_performance, + with two differences this population forces. Every executor that ever ran leaves a + terminal row behind, so an unfiltered "latest per scope" grows without bound here + in a way the controller one does not: this orders newest-first and takes a limit, + which puts the live executors -- the only ones still being snapshotted -- at the + top. And a closed executor's latest row IS its terminal row, so this answers "its + final value" and "its current value" with the same query. + + This reads the series, not memory: an executor younger than one snapshot interval + has no row yet, and a live one is up to an interval stale. That is deliberate -- + the last point of /performance/latest and the last point of /performance/history + are the same row. In-memory current figures are what /executors/ serves. + """ + latest = select( + ExecutorPerformanceSnapshot.executor_id, + func.max(ExecutorPerformanceSnapshot.timestamp).label("max_timestamp"), + ).group_by(ExecutorPerformanceSnapshot.executor_id) + + # The filters go on the grouped subquery rather than the join: they decide which + # executors are aggregated at all, instead of aggregating the whole table and + # throwing most of it away afterwards. Every one of them is constant across an + # executor's rows, so this cannot change which row wins the max. + if executor_id: + latest = latest.where(ExecutorPerformanceSnapshot.executor_id == executor_id) + if executor_type: + latest = latest.where(ExecutorPerformanceSnapshot.executor_type == executor_type) + if controller_id: + latest = latest.where(ExecutorPerformanceSnapshot.controller_id == controller_id) + if account_name: + latest = latest.where(ExecutorPerformanceSnapshot.account_name == account_name) + if connector_name: + latest = latest.where(ExecutorPerformanceSnapshot.connector_name == connector_name) + if trading_pair: + latest = latest.where(ExecutorPerformanceSnapshot.trading_pair == trading_pair) + latest = latest.subquery() + + query = ( + select(ExecutorPerformanceSnapshot) + .join( + latest, + (ExecutorPerformanceSnapshot.executor_id == latest.c.executor_id) & + (ExecutorPerformanceSnapshot.timestamp == latest.c.max_timestamp) + ) + .order_by(desc(ExecutorPerformanceSnapshot.timestamp)) + ) + if limit: + query = query.limit(limit) + + result = await self.session.execute(query) + return [self._to_dict(s) for s in result.scalars().all()] + + async def get_performance_history( + self, + executor_id: Optional[str] = None, + executor_type: Optional[str] = None, + controller_id: Optional[str] = None, + account_name: Optional[str] = None, + connector_name: Optional[str] = None, + trading_pair: Optional[str] = None, + limit: Optional[int] = None, + cursor: Optional[str] = None, + start_time: Optional[datetime] = None, + end_time: Optional[datetime] = None, + interval: str = "5m" + ) -> Tuple[List[Dict], Optional[str], bool]: + """Get a snapshot series with cursor pagination and interval sampling. + + `controller_id` filters within the executor population only. An in-process + executor's controller_id and a Docker bot's MQTT controller_id are not guaranteed + to name the same thing, so this is never a key to join the two subjects on. + """ + interval_minutes = self._interval_to_minutes(interval) + query = ( + select(ExecutorPerformanceSnapshot) + .order_by(desc(ExecutorPerformanceSnapshot.timestamp)) + ) + + if executor_id: + query = query.filter(ExecutorPerformanceSnapshot.executor_id == executor_id) + if executor_type: + query = query.filter(ExecutorPerformanceSnapshot.executor_type == executor_type) + if controller_id: + query = query.filter(ExecutorPerformanceSnapshot.controller_id == controller_id) + if account_name: + query = query.filter(ExecutorPerformanceSnapshot.account_name == account_name) + if connector_name: + query = query.filter(ExecutorPerformanceSnapshot.connector_name == connector_name) + if trading_pair: + query = query.filter(ExecutorPerformanceSnapshot.trading_pair == trading_pair) + if start_time: + query = query.filter(ExecutorPerformanceSnapshot.timestamp >= start_time) + if end_time: + query = query.filter(ExecutorPerformanceSnapshot.timestamp <= end_time) + if cursor: + try: + cursor_time = datetime.fromisoformat(cursor.replace('Z', '+00:00')) + query = query.filter(ExecutorPerformanceSnapshot.timestamp < cursor_time) + except (ValueError, TypeError): + pass + + # Over-fetch by the thinning ratio so a sampled page still fills up, plus one row + # to tell has_more apart from a page that happens to land exactly on the limit. + # The ratio is the worst case: sampling is per executor, so a window covering many + # executors retains more of what was fetched, never less. + sampling_multiplier = max(1, int(interval_minutes // self.grain_minutes)) + fetch_limit = (limit * sampling_multiplier + 1) if limit else (100 * sampling_multiplier + 1) + query = query.limit(fetch_limit) + + result = await self.session.execute(query) + snapshots = result.scalars().all() + + history = [self._to_dict(s) for s in snapshots] + + sampled = self._sample_by_interval(history, interval_minutes, self.grain_minutes) + + has_more = len(sampled) > limit if limit else False + if has_more: + sampled = sampled[:limit] + + next_cursor = None + if has_more and sampled: + next_cursor = sampled[-1]["timestamp"] + + return sampled, next_cursor, has_more + + async def prune_older_than(self, cutoff: datetime) -> Tuple[int, int]: + """Delete snapshots older than `cutoff` from both snapshot tables. + + Retention is one policy, not two: the operator sets how much performance history + to keep, and both series obey it. It lives here rather than on the controller + repository because that one is on the wire-compatible read path and this is where + the growth is generated. + + Returns (executor rows deleted, controller rows deleted). + """ + executor_result = await self.session.execute( + delete(ExecutorPerformanceSnapshot).where(ExecutorPerformanceSnapshot.timestamp < cutoff) + ) + controller_result = await self.session.execute( + delete(ControllerPerformanceSnapshot).where(ControllerPerformanceSnapshot.timestamp < cutoff) + ) + await self.session.flush() + return executor_result.rowcount or 0, controller_result.rowcount or 0 + + @staticmethod + def _to_dict(snapshot: ExecutorPerformanceSnapshot) -> Dict: + return { + "timestamp": snapshot.timestamp.isoformat(), + "executor_id": snapshot.executor_id, + "executor_type": snapshot.executor_type, + "account_name": snapshot.account_name, + "connector_name": snapshot.connector_name, + "trading_pair": snapshot.trading_pair, + "controller_id": snapshot.controller_id, + "status": snapshot.status, + "close_type": snapshot.close_type, + "is_terminal": bool(snapshot.is_terminal), + "net_pnl_quote": float(snapshot.net_pnl_quote or 0), + "net_pnl_pct": float(snapshot.net_pnl_pct or 0), + "cum_fees_quote": float(snapshot.cum_fees_quote or 0), + "filled_amount_quote": float(snapshot.filled_amount_quote or 0), + } diff --git a/database/repositories/executor_repository.py b/database/repositories/executor_repository.py index b434d9a0..e547712e 100644 --- a/database/repositories/executor_repository.py +++ b/database/repositories/executor_repository.py @@ -1,14 +1,20 @@ """ Repository for executor database operations. """ +import logging +import math from datetime import datetime, timezone from decimal import Decimal -from typing import Any, Dict, List, Optional +from typing import Any, Dict, List, Optional, Tuple from sqlalchemy import and_, case, desc, func, select +from sqlalchemy.exc import IntegrityError from sqlalchemy.ext.asyncio import AsyncSession from database.models import ExecutorOrder, ExecutorRecord, PositionHoldRecord +from database.repositories.executor_performance_repository import ExecutorPerformanceRepository + +logger = logging.getLogger(__name__) class ExecutorRepository: @@ -30,9 +36,15 @@ async def create_executor( trading_pair: str, config: Optional[str] = None, status: str = "RUNNING", - controller_id: str = "main" + controller_id: str = "main", + created_at: Optional[datetime] = None ) -> ExecutorRecord: - """Create a new executor record.""" + """Create a new executor record. + + `created_at` defaults to the server clock; pass it only when the row is being + written after the fact (see upsert_executor_completion) so the record still + orders by when the executor actually started. + """ executor = ExecutorRecord( executor_id=executor_id, executor_type=executor_type, @@ -43,6 +55,8 @@ async def create_executor( config=config, status=status ) + if created_at is not None: + executor.created_at = created_at self.session.add(executor) await self.session.flush() @@ -90,6 +104,75 @@ async def update_executor( return executor + async def upsert_executor_completion( + self, + executor_id: str, + executor_type: str, + account_name: str, + connector_name: str, + trading_pair: str, + controller_id: str = "main", + config: Optional[str] = None, + created_at: Optional[datetime] = None, + status: Optional[str] = None, + close_type: Optional[str] = None, + net_pnl_quote: Optional[Decimal] = None, + net_pnl_pct: Optional[Decimal] = None, + cum_fees_quote: Optional[Decimal] = None, + filled_amount_quote: Optional[Decimal] = None, + final_state: Optional[str] = None, + error_log: Optional[str] = None + ) -> Tuple[Optional[ExecutorRecord], bool]: + """Write an executor's final state, creating its row if it is not there yet. + + `update_executor` is select-then-update, so it silently does nothing when the + creation INSERT has not landed: an executor that closes milliseconds after + start can have its completion written before (or instead of) its creation row, + and the final state would just be dropped, leaving a phantom RUNNING executor. + This is the same insert-or-update shape `upsert_position_hold` already uses. + + Returns (record, created) where `created` is True if the row had to be + repaired — the caller is expected to log that, it is never normal. + """ + completion = dict( + status=status, + close_type=close_type, + net_pnl_quote=net_pnl_quote, + net_pnl_pct=net_pnl_pct, + cum_fees_quote=cum_fees_quote, + filled_amount_quote=filled_amount_quote, + final_state=final_state, + error_log=error_log, + ) + + executor = await self.update_executor(executor_id=executor_id, **completion) + if executor is not None: + return executor, False + + # No row: insert one from the metadata the caller carries. A creation INSERT + # racing us in another session can still win between our SELECT and this + # INSERT, so do it in a SAVEPOINT and fall back to the update — the outer + # transaction stays usable either way. + created = True + try: + async with self.session.begin_nested(): + await self.create_executor( + executor_id=executor_id, + executor_type=executor_type, + account_name=account_name, + connector_name=connector_name, + trading_pair=trading_pair, + config=config, + status=status or "RUNNING", + controller_id=controller_id, + created_at=created_at, + ) + except IntegrityError: + created = False + + executor = await self.update_executor(executor_id=executor_id, **completion) + return executor, created + async def get_executor_by_id(self, executor_id: str) -> Optional[ExecutorRecord]: """Get an executor by ID.""" stmt = select(ExecutorRecord).where(ExecutorRecord.executor_id == executor_id) @@ -316,58 +399,50 @@ async def clear_position_hold( return result.rowcount > 0 async def get_executor_stats(self) -> Dict[str, Any]: - """Get statistics about executors.""" - # Total executors - total_stmt = select(func.count(ExecutorRecord.id)) - total_result = await self.session.execute(total_stmt) - total_executors = total_result.scalar() or 0 - - # Active executors - active_stmt = select(func.count(ExecutorRecord.id)).where( - ExecutorRecord.status == "RUNNING" - ) - active_result = await self.session.execute(active_stmt) - active_executors = active_result.scalar() or 0 - - # Total PnL - pnl_stmt = select(func.sum(ExecutorRecord.net_pnl_quote)) - pnl_result = await self.session.execute(pnl_stmt) - total_pnl = pnl_result.scalar() or Decimal("0") + """Get statistics about executors. - # Total volume — the volume generated, not the capital deployed. - volume_stmt = select(func.sum(ExecutorRecord.filled_amount_quote)) - volume_result = await self.session.execute(volume_stmt) - total_volume = volume_result.scalar() or Decimal("0") + Two round-trips, not seven: one row of scalar aggregates, and one grouped + statement the three per-column breakdowns are pivoted out of in Python. The + aggregate shape is the same one `get_performance_report` uses next door. + """ + agg_stmt = select( + func.count(ExecutorRecord.id).label("total"), + func.coalesce(func.sum(case( + (ExecutorRecord.status == "RUNNING", 1), + else_=0, + )), 0).label("active"), + func.coalesce(func.sum(ExecutorRecord.net_pnl_quote), Decimal(0)).label("pnl"), + # The volume generated, not the capital deployed. + func.coalesce(func.sum(ExecutorRecord.filled_amount_quote), Decimal(0)).label("vol"), + ) + agg_row = (await self.session.execute(agg_stmt)).one() - # Executors by type - type_stmt = select( + # One grouped statement carries all three breakdowns. The number of rows it + # returns is bounded by the distinct (type, status, connector) combinations -- + # a handful in practice -- and each breakdown is summed back out of them below. + group_stmt = select( + ExecutorRecord.executor_type, + ExecutorRecord.status, + ExecutorRecord.connector_name, + func.count(ExecutorRecord.id).label("count"), + ).group_by( ExecutorRecord.executor_type, - func.count(ExecutorRecord.id).label('count') - ).group_by(ExecutorRecord.executor_type) - type_result = await self.session.execute(type_stmt) - type_counts = {row.executor_type: row.count for row in type_result} - - # Executors by status - status_stmt = select( ExecutorRecord.status, - func.count(ExecutorRecord.id).label('count') - ).group_by(ExecutorRecord.status) - status_result = await self.session.execute(status_stmt) - status_counts = {row.status: row.count for row in status_result} - - # Executors by connector - connector_stmt = select( ExecutorRecord.connector_name, - func.count(ExecutorRecord.id).label('count') - ).group_by(ExecutorRecord.connector_name) - connector_result = await self.session.execute(connector_stmt) - connector_counts = {row.connector_name: row.count for row in connector_result} + ) + type_counts: Dict[str, int] = {} + status_counts: Dict[str, int] = {} + connector_counts: Dict[str, int] = {} + for row in await self.session.execute(group_stmt): + type_counts[row.executor_type] = type_counts.get(row.executor_type, 0) + row.count + status_counts[row.status] = status_counts.get(row.status, 0) + row.count + connector_counts[row.connector_name] = connector_counts.get(row.connector_name, 0) + row.count return { - "total_executors": total_executors, - "active_executors": active_executors, - "total_pnl_quote": float(total_pnl), - "total_volume_quote": float(total_volume), + "total_executors": agg_row.total or 0, + "active_executors": agg_row.active or 0, + "total_pnl_quote": float(agg_row.pnl), + "total_volume_quote": float(agg_row.vol), "type_counts": type_counts, "status_counts": status_counts, "connector_counts": connector_counts @@ -380,7 +455,11 @@ async def get_performance_report( """Get a performance report, optionally filtered by controller_id. Returns aggregate metrics: total executors, PnL, fees, volume, - win rate, per-executor PnL list (for Sharpe), and breakdown by type. + win rate, PnL dispersion (for Sharpe), and breakdown by type. + + Every metric is an aggregate: the number of rows this returns does not + grow with the size of the executors table. It is polled by the + `/ws/executors` performance push loop, once per interval per subscriber. """ base_filter = [] if controller_id: @@ -408,6 +487,10 @@ async def get_performance_report( func.coalesce(func.sum(ExecutorRecord.cum_fees_quote), Decimal(0)).label("fees"), func.coalesce(func.sum(ExecutorRecord.filled_amount_quote), Decimal(0)).label("vol"), func.coalesce(func.avg(ExecutorRecord.net_pnl_pct), Decimal(0)).label("pnl_pct_avg"), + func.coalesce( + func.sum(ExecutorRecord.net_pnl_quote * ExecutorRecord.net_pnl_quote), + Decimal(0), + ).label("pnl_sq"), func.count(ExecutorRecord.id).label("completed_count"), func.sum(case( (ExecutorRecord.net_pnl_quote > 0, 1), @@ -420,12 +503,19 @@ async def get_performance_report( wins = agg_row.wins or 0 win_rate = (wins / completed_count) if completed_count > 0 else 0.0 - # --- Per-executor PnL list for Sharpe (excluding POSITION_HOLD) --- - pnl_list_stmt = select(ExecutorRecord.net_pnl_quote).where( - and_(*completed_filter) - ) - pnl_rows = await self.session.execute(pnl_list_stmt) - pnl_values = [float(r[0] or 0) for r in pnl_rows] + # --- PnL dispersion for Sharpe, from the aggregates rather than the rows --- + # Sample standard deviation out of the count, the sum and the sum of squares the + # query above already carries. Fetching one net_pnl_quote per completed executor + # to do this in Python cost a full table scan on every poll of this report. + # NULL PnL counts as zero, as the old per-row list did: SUM skips it, COUNT does not. + # The moments stay Decimal until the square root: a large mean with a small spread + # loses its whole variance to cancellation if the subtraction is done in float. + pnl_std = None + if completed_count >= 2: + variance = ( + agg_row.pnl_sq - agg_row.pnl * agg_row.pnl / completed_count + ) / (completed_count - 1) + pnl_std = math.sqrt(max(float(variance), 0.0)) # max() clamps noise at zero variance # --- Breakdown by executor type (also excluding POSITION_HOLD to match aggregate totals) --- type_stmt = select( @@ -467,7 +557,8 @@ async def get_performance_report( "fees_total_quote": float(agg_row.fees), "volume_total_quote": float(agg_row.vol), "win_rate": win_rate, - "pnl_values": pnl_values, + "completed_count": completed_count, + "pnl_std": pnl_std, "by_type": by_type, } @@ -560,38 +651,105 @@ async def cleanup_orphaned_executors( ) -> int: """ Clean up orphaned executors - those marked as RUNNING but not in active memory. + + Each terminated row takes the metrics of its most recent performance snapshot. + A running executor's row carries the zeros it was INSERTed with -- the live + figures only exist in ExecutorService memory until it completes -- so terminating + it without them booked every executor that was live at a restart at 0 PnL, 0 fees + and 0 volume, permanently, and /executors/performance summed those zeros forever. + The snapshot series is what makes the real figures recoverable here. + + An executor with no snapshot (created and orphaned inside one snapshot interval) + keeps the old behaviour: there is nothing better to write. The adopted figures are + the last observed ones rather than the true final ones, which is why the row is + still marked SYSTEM_CLEANUP -- a reader can tell an approximated close from a + clean one. + + Each reap also writes the terminal snapshot row that the normal completion path + writes, in this same transaction. Without it the reap was the one way an executor + could reach TERMINATED with no terminal row behind it, and the invariant the + /performance routes are built on -- a closed executor's latest row IS its terminal + row, answerable with no join back to `executors` -- held for every executor except + the ones a crash caught. /performance/latest served their last RUNNING snapshot + forever, is_terminal false and close_type null, while /executors/{id} reported + TERMINATED/SYSTEM_CLEANUP: two surfaces permanently disagreeing about whether an + executor was done. + + The row carries the record's post-adoption figures, so it says exactly what the + record says, and its SYSTEM_CLEANUP close_type marks the series end as an + approximated close the same way the record does. + Args: active_executor_ids: List of executor IDs currently active in memory close_type: Close type to set for cleaned up executors Returns: Number of executors cleaned up """ - from sqlalchemy import update - # Find executors that are RUNNING but not in the active list conditions = [ExecutorRecord.status == "RUNNING"] if active_executor_ids: conditions.append(~ExecutorRecord.executor_id.in_(active_executor_ids)) - # First, get the count of orphaned executors for logging - count_stmt = select(func.count(ExecutorRecord.id)).where(and_(*conditions)) - count_result = await self.session.execute(count_stmt) - orphaned_count = count_result.scalar() or 0 - - if orphaned_count > 0: - # Update orphaned executors to TERMINATED status - update_stmt = ( - update(ExecutorRecord) - .where(and_(*conditions)) - .values( - status="TERMINATED", - close_type=close_type, - closed_at=datetime.now(timezone.utc) - ) - ) + orphaned = (await self.session.execute( + select(ExecutorRecord).where(and_(*conditions)) + )).scalars().all() - await self.session.execute(update_stmt) - await self.session.flush() + if not orphaned: + return 0 + + latest_snapshots = await ExecutorPerformanceRepository(self.session).get_latest_for( + [record.executor_id for record in orphaned] + ) + + closed_at = datetime.now(timezone.utc) + terminal_rows = [] + for record in orphaned: + record.status = "TERMINATED" + record.close_type = close_type + record.closed_at = closed_at + + snapshot = latest_snapshots.get(record.executor_id) + if snapshot: + record.net_pnl_quote = Decimal(str(snapshot["net_pnl_quote"])) + record.net_pnl_pct = Decimal(str(snapshot["net_pnl_pct"])) + record.cum_fees_quote = Decimal(str(snapshot["cum_fees_quote"])) + record.filled_amount_quote = Decimal(str(snapshot["filled_amount_quote"])) + + terminal_rows.append({ + "executor_id": record.executor_id, + # The identity comes from the record, not the snapshot: an executor + # orphaned before its first snapshot has no snapshot to read it from, and + # these columns are constant across an executor's life either way. + "executor_type": record.executor_type or "unknown", + "account_name": record.account_name, + "connector_name": record.connector_name or "", + "trading_pair": record.trading_pair or "", + "controller_id": record.controller_id or "main", + "status": "TERMINATED", + "close_type": close_type, + "is_terminal": True, + "net_pnl_quote": record.net_pnl_quote, + "net_pnl_pct": record.net_pnl_pct, + "cum_fees_quote": record.cum_fees_quote, + "filled_amount_quote": record.filled_amount_quote, + "snapshot_timestamp": closed_at, + }) + + await self.session.flush() + + # Behind a SAVEPOINT for the same reason _persist_executor_completed's terminal + # row is: the reap is the accounting and the snapshot is a point on a chart, so a + # failing INSERT here must roll back only itself rather than abort the reap and + # leave the whole fleet RUNNING. + try: + async with self.session.begin_nested(): + await ExecutorPerformanceRepository(self.session).save_snapshots(terminal_rows) + except Exception as e: + logger.error( + f"Reaped {len(orphaned)} orphaned executors but could not write their " + f"terminal performance snapshots; their series end at the last periodic " + f"row: {e}" + ) - return orphaned_count + return len(orphaned) diff --git a/database/repositories/gateway_clmm_repository.py b/database/repositories/gateway_clmm_repository.py index 2861fb3b..3e599818 100644 --- a/database/repositories/gateway_clmm_repository.py +++ b/database/repositories/gateway_clmm_repository.py @@ -2,11 +2,14 @@ from decimal import Decimal from typing import Dict, List, Optional, Set -from sqlalchemy import select +from sqlalchemy import or_, select from sqlalchemy.ext.asyncio import AsyncSession from database.models import GatewayCLMMEvent, GatewayCLMMPosition +# The event types whose gas_token the pre-ARCH-054 write paths damaged. +LIQUIDITY_EVENT_TYPES = ("ADD_LIQUIDITY", "REMOVE_LIQUIDITY") + class GatewayCLMMRepository: def __init__(self, session: AsyncSession): @@ -371,6 +374,66 @@ async def get_position_events( result = await self.session.execute(query) return result.scalars().all() + async def backfill_liquidity_gas_tokens(self) -> Dict: + """Repair gas_token on liquidity events written before ARCH-054 unified the chain map. + + Two write paths damaged these rows and neither will revisit them: the add/remove + handlers' two-branch ternary wrote gas_token NULL off solana/ethereum, and a row + inserted CONFIRMED is never re-polled, so that NULL is permanent; the poller's old + 6-entry dict wrote "UNKNOWN" for base, arbitrum and polygon. The write paths are + fixed, the historical rows are not, and every cost report reading them is wrong. + + Only rows that actually recorded a gas fee are repaired: with no fee there is no + currency to name, which is exactly what the routes persist today + (``get_native_gas_token(chain) if gas_fee is not None else None``). The chain comes + from the position's ``network`` ("solana-mainnet-beta" -> "solana") and is resolved + through ``get_native_gas_token`` — the single source ARCH-054 established, never a + second copy of the map. A chain that still resolves to "UNKNOWN" is left untouched + and counted, so an unmapped chain surfaces instead of being papered over. + + Idempotent: a repaired row no longer matches the NULL/"UNKNOWN" filter, so a second + run fixes nothing and changes nothing. + + Returns: + {"fixed": int, "unresolved": int, "unresolved_networks": List[str]} + """ + # Imported here, not at module scope: `services/__init__` imports back into + # `database`, so a top-level import of it from a repository is a cycle. + from services.gateway_client import get_native_gas_token + + query = ( + select(GatewayCLMMEvent, GatewayCLMMPosition.network) + .join(GatewayCLMMPosition, GatewayCLMMEvent.position_id == GatewayCLMMPosition.id) + .where( + GatewayCLMMEvent.event_type.in_(LIQUIDITY_EVENT_TYPES), + GatewayCLMMEvent.gas_fee.isnot(None), + or_(GatewayCLMMEvent.gas_token.is_(None), GatewayCLMMEvent.gas_token == "UNKNOWN"), + ) + ) + result = await self.session.execute(query) + + fixed = 0 + unresolved = 0 + unresolved_networks: Set[str] = set() + for event, network in result.all(): + chain = (network or "").split("-", 1)[0] + gas_token = get_native_gas_token(chain) + if gas_token == "UNKNOWN": + unresolved += 1 + unresolved_networks.add(network) + continue + event.gas_token = gas_token + fixed += 1 + + if fixed: + await self.session.flush() + + return { + "fixed": fixed, + "unresolved": unresolved, + "unresolved_networks": sorted(unresolved_networks), + } + async def get_pending_events(self, limit: int = 100) -> List[GatewayCLMMEvent]: """Get events that are still pending confirmation.""" query = select(GatewayCLMMEvent).where( diff --git a/database/repositories/order_repository.py b/database/repositories/order_repository.py index 5036bb10..49b38afa 100644 --- a/database/repositories/order_repository.py +++ b/database/repositories/order_repository.py @@ -1,3 +1,4 @@ +import logging from datetime import datetime from decimal import Decimal from typing import Dict, List, Optional @@ -7,8 +8,18 @@ from database.models import Order +logger = logging.getLogger(__name__) + class OrderRepository: + # Client order ids are batched into `IN (...)` clauses of at most this size so a + # connector with a very large book does not build an unbounded bind-parameter list. + CLIENT_ID_CHUNK_SIZE = 500 + + # Default cap for `get_active_orders`. Callers that need the complete book + # (connector startup, for instance) must pass `limit=None`. + DEFAULT_ACTIVE_ORDERS_LIMIT = 1000 + def __init__(self, session: AsyncSession): self.session = session @@ -26,6 +37,27 @@ async def get_order_by_client_id(self, client_order_id: str) -> Optional[Order]: ) return result.scalar_one_or_none() + async def get_orders_by_client_ids(self, client_order_ids: List[str]) -> List[Order]: + """Get the orders matching a batch of client order IDs. + + Batched sibling of `get_order_by_client_id`: one query per `CLIENT_ID_CHUNK_SIZE` + ids instead of one round trip per id. Ids with no row are simply absent from the + result; no status filter is applied, so rows already in a terminal state come back + too and can still be corrected. + """ + if not client_order_ids: + return [] + + orders: List[Order] = [] + ids = list(client_order_ids) + for start in range(0, len(ids), self.CLIENT_ID_CHUNK_SIZE): + chunk = ids[start:start + self.CLIENT_ID_CHUNK_SIZE] + result = await self.session.execute( + select(Order).where(Order.client_order_id.in_(chunk)) + ) + orders.extend(result.scalars().all()) + return orders + async def update_order_status(self, client_order_id: str, status: str, error_message: Optional[str] = None) -> Optional[Order]: """Update order status and optional error message.""" @@ -38,8 +70,8 @@ async def update_order_status(self, client_order_id: str, status: str, return order async def update_order_fill(self, client_order_id: str, filled_amount: Decimal, - average_fill_price: Decimal, fee_paid: Decimal = None, - fee_currency: str = None, exchange_order_id: str = None) -> Optional[Order]: + average_fill_price: Decimal, fee_paid: Decimal = None, + fee_currency: str = None, exchange_order_id: str = None) -> Optional[Order]: """Update order with fill information.""" result = await self.session.execute( select(Order).where(Order.client_order_id == client_order_id) @@ -49,10 +81,10 @@ async def update_order_fill(self, client_order_id: str, filled_amount: Decimal, # Add to existing filled amount instead of replacing previous_filled = Decimal(str(order.filled_amount or 0)) order.filled_amount = float(previous_filled + filled_amount) - + # Update average price (simplified - use latest fill price) order.average_fill_price = float(average_fill_price) - + # Add to existing fees if fee_paid is not None: previous_fee = Decimal(str(order.fee_paid or 0)) @@ -61,27 +93,27 @@ async def update_order_fill(self, client_order_id: str, filled_amount: Decimal, order.fee_currency = fee_currency if exchange_order_id: order.exchange_order_id = exchange_order_id - + # Update status based on total filled amount total_filled = Decimal(str(order.filled_amount)) if total_filled >= Decimal(str(order.amount)): order.status = "FILLED" elif total_filled > 0: order.status = "PARTIALLY_FILLED" - + await self.session.flush() return order - async def get_orders(self, account_name: Optional[str] = None, - connector_name: Optional[str] = None, - trading_pair: Optional[str] = None, - status: Optional[str] = None, - start_time: Optional[int] = None, - end_time: Optional[int] = None, - limit: int = 100, offset: int = 0) -> List[Order]: + async def get_orders(self, account_name: Optional[str] = None, + connector_name: Optional[str] = None, + trading_pair: Optional[str] = None, + status: Optional[str] = None, + start_time: Optional[int] = None, + end_time: Optional[int] = None, + limit: int = 100, offset: int = 0) -> List[Order]: """Get orders with filtering and pagination.""" query = select(Order) - + # Apply filters if account_name: query = query.where(Order.account_name == account_name) @@ -97,22 +129,28 @@ async def get_orders(self, account_name: Optional[str] = None, if end_time: end_dt = datetime.fromtimestamp(end_time / 1000) query = query.where(Order.created_at <= end_dt) - + # Apply ordering and pagination query = query.order_by(Order.created_at.desc()) query = query.limit(limit).offset(offset) - + result = await self.session.execute(query) return result.scalars().all() async def get_active_orders(self, account_name: Optional[str] = None, - connector_name: Optional[str] = None, - trading_pair: Optional[str] = None) -> List[Order]: - """Get active orders (SUBMITTED, OPEN, PARTIALLY_FILLED, PENDING_CANCEL).""" + connector_name: Optional[str] = None, + trading_pair: Optional[str] = None, + limit: Optional[int] = DEFAULT_ACTIVE_ORDERS_LIMIT) -> List[Order]: + """Get active orders (SUBMITTED, OPEN, PARTIALLY_FILLED, PENDING_CANCEL). + + `limit` caps how many rows come back, newest first; pass `limit=None` to get the + whole book with no cap. A query that reaches the cap logs a warning, so a + truncated result is never silent. + """ query = select(Order).where( Order.status.in_(["SUBMITTED", "OPEN", "PARTIALLY_FILLED", "PENDING_CANCEL"]) ) - + # Apply filters if account_name: query = query.where(Order.account_name == account_name) @@ -120,15 +158,26 @@ async def get_active_orders(self, account_name: Optional[str] = None, query = query.where(Order.connector_name == connector_name) if trading_pair: query = query.where(Order.trading_pair == trading_pair) - - query = query.order_by(Order.created_at.desc()).limit(1000) - + + query = query.order_by(Order.created_at.desc()) + if limit is not None: + query = query.limit(limit) + result = await self.session.execute(query) - return result.scalars().all() + orders = result.scalars().all() + + if limit is not None and len(orders) >= limit: + logger.warning( + f"get_active_orders returned {len(orders)} orders, reaching its limit of {limit}; " + f"older active orders may have been truncated " + f"(account={account_name}, connector={connector_name}, trading_pair={trading_pair})" + ) + + return orders async def get_orders_summary(self, account_name: Optional[str] = None, - start_time: Optional[int] = None, - end_time: Optional[int] = None) -> Dict: + start_time: Optional[int] = None, + end_time: Optional[int] = None) -> Dict: """Get order summary statistics using a single DB-level aggregate query.""" query = select(Order.status, func.count()).group_by(Order.status) @@ -182,4 +231,4 @@ def to_dict(self, order: Order) -> Dict: "updated_at": order.updated_at.isoformat(), "exchange_order_id": order.exchange_order_id, "error_message": order.error_message, - } \ No newline at end of file + } diff --git a/deps.py b/deps.py index 7b4ff4ae..bb4dc640 100644 --- a/deps.py +++ b/deps.py @@ -1,4 +1,6 @@ -from fastapi import Request +import math + +from fastapi import Depends, HTTPException, Request from database import AsyncDatabaseManager from services.accounts_service import AccountsService @@ -7,7 +9,11 @@ from services.docker_service import DockerService from services.executor_service import ExecutorService from services.executor_ws_manager import ExecutorWebSocketManager +from services.gateway_amm_service import GatewayAMMService +from services.gateway_client import GatewayClient +from services.gateway_clmm_service import GatewayCLMMService from services.gateway_service import GatewayService +from services.gateway_swap_service import GatewaySwapService from services.market_data_service import MarketDataService from services.trading_history_service import TradingHistoryService from services.trading_service import TradingService @@ -26,6 +32,30 @@ def get_accounts_service(request: Request) -> AccountsService: return request.app.state.accounts_service +GATEWAY_UNAVAILABLE_DETAIL = "Gateway service is not available" +# The guard's verdict is cached for the client's ping TTL, so a retry sooner than that can only +# read back the same answer. Advertise the TTL as the standard back-off signal. +GATEWAY_UNAVAILABLE_RETRY_AFTER = str(math.ceil(GatewayClient.PING_CACHE_TTL_SECONDS)) + + +async def require_gateway_online(accounts_service: AccountsService = Depends(get_accounts_service)) -> None: + """Pre-flight guard for routes that cannot be served at all while Gateway is unreachable. + + Attach per route with ``dependencies=[Depends(require_gateway_online)]``. Deliberately not applied at + router level: the container-control routes (``/gateway/start``, ``/gateway/restart``, ...) must answer + precisely when Gateway is down, and the DB-backed reads keep serving stored data regardless. + + ``ping`` answers from a short-lived cache, so a burst of guarded requests costs one round-trip + rather than one each; an unreachable Gateway clears that cache and is reported straight away. + """ + if not await accounts_service.gateway_client.ping(): + raise HTTPException( + status_code=503, + detail=GATEWAY_UNAVAILABLE_DETAIL, + headers={"Retry-After": GATEWAY_UNAVAILABLE_RETRY_AFTER}, + ) + + def get_docker_service(request: Request) -> DockerService: """Get DockerService from app state.""" return request.app.state.docker_service @@ -36,6 +66,21 @@ def get_gateway_service(request: Request) -> GatewayService: return request.app.state.gateway_service +def get_gateway_clmm_service(request: Request) -> GatewayCLMMService: + """Get GatewayCLMMService from app state.""" + return request.app.state.gateway_clmm_service + + +def get_gateway_amm_service(request: Request) -> GatewayAMMService: + """Get GatewayAMMService from app state.""" + return request.app.state.gateway_amm_service + + +def get_gateway_swap_service(request: Request) -> GatewaySwapService: + """Get GatewaySwapService from app state.""" + return request.app.state.gateway_swap_service + + def get_connector_service(request: Request) -> UnifiedConnectorService: """Get UnifiedConnectorService from app state.""" return request.app.state.connector_service diff --git a/environment.yml b/environment.yml index b131754d..4fb7d728 100644 --- a/environment.yml +++ b/environment.yml @@ -32,6 +32,8 @@ dependencies: - msgpack>=1.0.5 - flake8 - isort + - pytest + - pytest-asyncio - pre-commit - logfire - logfire[fastapi] diff --git a/main.py b/main.py index bc493f8c..c87d226a 100644 --- a/main.py +++ b/main.py @@ -54,6 +54,7 @@ def patched_save_to_yml(yml_path, cm): gateway_clmm, gateway_swap, market_data, + performance, portfolio, scripts, storage, @@ -67,7 +68,10 @@ def patched_save_to_yml(yml_path, cm): from services.docker_service import DockerService # noqa: E402 from services.executor_service import ExecutorService # noqa: E402 from services.executor_ws_manager import ExecutorWebSocketManager # noqa: E402 +from services.gateway_amm_service import GatewayAMMService # noqa: E402 +from services.gateway_clmm_service import GatewayCLMMService # noqa: E402 from services.gateway_service import GatewayService # noqa: E402 +from services.gateway_swap_service import GatewaySwapService # noqa: E402 from services.market_data_service import MarketDataService # noqa: E402 from services.trading_history_service import TradingHistoryService # noqa: E402 from services.trading_service import TradingService # noqa: E402 @@ -237,6 +241,13 @@ async def lifespan(app: FastAPI): trading_history_service = TradingHistoryService(db_manager=db_manager) logging.info("TradingHistoryService initialized") + # Gateway CLMM/swap persistence - the /gateway/clmm/* and /gateway/swap* routes + # read and write their history through these instead of owning sessions themselves + gateway_amm_service = GatewayAMMService(db_manager=db_manager) + gateway_clmm_service = GatewayCLMMService(db_manager=db_manager) + gateway_swap_service = GatewaySwapService(db_manager=db_manager) + logging.info("GatewayCLMMService and GatewaySwapService initialized") + # ========================================================================= # 4. ExecutorService - depends on TradingService (NO circular dependency) # ========================================================================= @@ -246,7 +257,9 @@ async def lifespan(app: FastAPI): db_manager=db_manager, default_account="master_account", update_interval=1.0, - max_retries=10 + max_retries=10, + performance_snapshot_interval=settings.performance.executor_snapshot_interval, + performance_retention_days=settings.performance.retention_days ) logging.info("ExecutorService initialized") @@ -331,6 +344,9 @@ async def lifespan(app: FastAPI): app.state.trading_service = trading_service app.state.accounts_service = accounts_service app.state.trading_history_service = trading_history_service + app.state.gateway_amm_service = gateway_amm_service + app.state.gateway_clmm_service = gateway_clmm_service + app.state.gateway_swap_service = gateway_swap_service app.state.executor_service = executor_service websocket_manager = WebSocketManager(market_data_service) app.state.websocket_manager = websocket_manager @@ -460,6 +476,7 @@ def auth_user( app.include_router(controllers.router, dependencies=[Depends(auth_user)]) app.include_router(scripts.router, dependencies=[Depends(auth_user)]) app.include_router(market_data.router, dependencies=[Depends(auth_user)]) +app.include_router(performance.router, dependencies=[Depends(auth_user)]) app.include_router(backtesting.router, dependencies=[Depends(auth_user)]) app.include_router(archived_bots.router, dependencies=[Depends(auth_user)]) app.include_router(storage.router, dependencies=[Depends(auth_user)]) diff --git a/models/bot_orchestration.py b/models/bot_orchestration.py index 5378ef0b..110a76b8 100644 --- a/models/bot_orchestration.py +++ b/models/bot_orchestration.py @@ -18,7 +18,7 @@ def _validate_safe_name(name: str, label: str) -> str: return name -def _validate_safe_config_name(name: str, label: str) -> str: +def validate_safe_config_name(name: str, label: str) -> str: """Validate a config file name, ignoring an optional .yml extension before checking the base name.""" base_name = name[:-4] if name.endswith(".yml") else name _validate_safe_name(base_name, label) @@ -140,7 +140,7 @@ def _validate_credentials_profile(cls, v: str) -> str: def _validate_script_config(cls, v: Optional[str]) -> Optional[str]: if v is None: return v - return _validate_safe_config_name(v, "script_config") + return validate_safe_config_name(v, "script_config") class V2ControllerDeployment(BaseModel): @@ -173,11 +173,11 @@ def _validate_credentials_profile(cls, v: str) -> str: @field_validator("controllers_config") @classmethod def _validate_controllers_config(cls, v: List[str]) -> List[str]: - return [_validate_safe_config_name(controller, "controllers_config") for controller in v] + return [validate_safe_config_name(controller, "controllers_config") for controller in v] @field_validator("script_config") @classmethod def _validate_script_config(cls, v: Optional[str]) -> Optional[str]: if v is None: return v - return _validate_safe_config_name(v, "script_config") + return validate_safe_config_name(v, "script_config") diff --git a/models/executors.py b/models/executors.py index 86287659..f5cc9b24 100644 --- a/models/executors.py +++ b/models/executors.py @@ -4,7 +4,7 @@ These models wrap Hummingbot's executor configuration types and provide validation for the REST API. """ -from datetime import datetime +from datetime import datetime, timezone from decimal import Decimal from typing import Any, Dict, List, Literal, Optional @@ -128,7 +128,7 @@ def add_fill( if executor_id and executor_id not in self.executor_ids: self.executor_ids.append(executor_id) - self.last_updated = datetime.utcnow() + self.last_updated = datetime.now(timezone.utc) def _calculate_realized_pnl(self): """Calculate realized PnL from matched buy/sell pairs and settle matched volume. @@ -185,7 +185,7 @@ def merge(self, other: "PositionHold"): self.executor_ids.append(eid) self._calculate_realized_pnl() - self.last_updated = datetime.utcnow() + self.last_updated = datetime.now(timezone.utc) class PositionHoldResponse(BaseModel): diff --git a/models/performance.py b/models/performance.py new file mode 100644 index 00000000..fc97939f --- /dev/null +++ b/models/performance.py @@ -0,0 +1,175 @@ +"""One normalized performance row, for both the controller and the executor series. + +The two populations are stored differently on purpose -- the controller's payload is an +opaque, core-versioned PerformanceReport that has to be blobbed, the executor's is +ExecutorInfo, whose fields this repo names all over the place -- but what a consumer +needs symmetric is the *response*, not the storage. So both are mapped into the shape +below and a client writes seriesFor(scope) once instead of two clients. + +The normalization is additive, never lossy: everything the controller report carries that +has no executor counterpart (open_order_volume, inventory_imbalance, positions_summary, +close_type_counts) stays reachable in the `performance` passthrough. +""" +from typing import Any, Dict, List, Optional + +from pydantic import BaseModel, Field + +SUBJECT_CONTROLLER = "controller" +SUBJECT_EXECUTOR = "executor" + +# "Settled" for an executor means its position is really closed. A POSITION_HOLD close +# hands the position on to position_holds, so counting it as realized here would +# double-count it -- the same exclusion ExecutorRepository.get_performance_report applies. +_UNSETTLED_CLOSE_TYPE = "POSITION_HOLD" + + +class PerformanceRow(BaseModel): + """A single point of a performance series, for either subject.""" + + timestamp: str = Field(description="ISO 8601 timestamp of the snapshot") + subject: str = Field(description='Which population this row came from: "controller" or "executor"') + scope_id: str = Field(description="controller_id for controllers, executor_id for executors") + status: str = Field(description="Status at the moment of the snapshot") + is_terminal: bool = Field( + default=False, + description="True only on the final row of a completed executor's series. " + "Always false for controllers, which have no terminal row." + ) + + realized_pnl_quote: float = Field( + description="Realized PnL in quote currency. For an executor this is its net PnL " + "once the position is really settled, and 0 while it is not. A " + "POSITION_HOLD close is not settled: the position carries on in " + "position_holds, so a held executor's PnL stays unrealized here even " + "though is_terminal is true." + ) + unrealized_pnl_quote: float = Field( + description="Unrealized PnL in quote currency. For an executor this is its net " + "PnL while the position is open, and after a POSITION_HOLD close." + ) + global_pnl_quote: float = Field(description="Total PnL (realized + unrealized) in quote currency") + global_pnl_pct: float = Field(description="Total PnL as a fraction of the capital deployed") + volume_quote: float = Field( + description="Volume traded in quote currency. Volume GENERATED, not capital " + "deployed: an LP position's deposit is excluded." + ) + cum_fees_quote: Optional[float] = Field( + default=None, + description="Cumulative fees in quote currency. NULL for controllers: their " + "PerformanceReport genuinely has no fees field, and unknown is not zero." + ) + + bot_name: Optional[str] = Field(default=None, description="Docker bot name; controllers only") + controller_id: Optional[str] = Field(default=None, description="Controller ID, on both subjects") + executor_id: Optional[str] = Field(default=None, description="Executor ID; executors only") + executor_type: Optional[str] = Field(default=None, description="Executor type; executors only") + account_name: Optional[str] = Field(default=None, description="Account name; executors only") + connector_name: Optional[str] = Field(default=None, description="Connector name; executors only") + trading_pair: Optional[str] = Field(default=None, description="Trading pair; executors only") + close_type: Optional[str] = Field(default=None, description="Close type; set on an executor's terminal row") + + performance: Dict[str, Any] = Field( + default_factory=dict, + description="Controllers: the raw stored PerformanceReport. Empty for executors, " + "whose metrics are all typed columns above." + ) + custom_info: Dict[str, Any] = Field( + default_factory=dict, + description="Controllers: the raw stored custom_info. Empty for executors -- these " + "payloads carry fill_events and grid levels, which have no business in " + "a per-minute row." + ) + + +class PerformanceHistoryResponse(BaseModel): + """Envelope of GET /performance/history, matching the controller route's shape.""" + + status: str = Field(default="success") + data: List[PerformanceRow] + pagination: Dict[str, Any] = Field( + description="next_cursor, has_more, limit and the interval that was requested" + ) + + +class PerformanceLatestResponse(BaseModel): + """Envelope of GET /performance/latest. + + No pagination block: this is one row per scope, not a series, so there is no cursor + to walk and no interval to sample. `limit` is a cap on how many scopes come back, + newest-first, not a page boundary. + """ + + status: str = Field(default="success") + data: List[PerformanceRow] + + +def _as_float(value: Any) -> float: + try: + return float(value) + except (TypeError, ValueError): + return 0.0 + + +def controller_row_to_performance_row(row: Dict[str, Any]) -> PerformanceRow: + """Map a ControllerPerformanceRepository row into the normalized shape. + + The report is read by name but never required: a controller that reported nothing + yields a row of zeros with its raw payload still attached, rather than disappearing + from the series. + """ + performance = row.get("performance") or {} + + return PerformanceRow( + timestamp=row["timestamp"], + subject=SUBJECT_CONTROLLER, + scope_id=row.get("controller_id") or "", + status=row.get("status") or "unknown", + is_terminal=False, + realized_pnl_quote=_as_float(performance.get("realized_pnl_quote")), + unrealized_pnl_quote=_as_float(performance.get("unrealized_pnl_quote")), + global_pnl_quote=_as_float(performance.get("global_pnl_quote")), + global_pnl_pct=_as_float(performance.get("global_pnl_pct")), + volume_quote=_as_float(performance.get("volume_traded")), + # Deliberately not 0.0: PerformanceReport has no fees field, and a consumer + # charting fees has to be able to tell "not measured" from "measured and empty". + cum_fees_quote=None, + bot_name=row.get("bot_name"), + controller_id=row.get("controller_id"), + performance=performance, + custom_info=row.get("custom_info") or {}, + ) + + +def executor_row_to_performance_row(row: Dict[str, Any]) -> PerformanceRow: + """Map an ExecutorPerformanceRepository row into the normalized shape. + + An executor reports one net PnL, not a realized/unrealized pair, so the split is made + from whether the position is settled: everything is unrealized while it is open, and + realized once it closes for real. + """ + net_pnl_quote = _as_float(row.get("net_pnl_quote")) + settled = bool(row.get("is_terminal")) and row.get("close_type") != _UNSETTLED_CLOSE_TYPE + + return PerformanceRow( + timestamp=row["timestamp"], + subject=SUBJECT_EXECUTOR, + scope_id=row.get("executor_id") or "", + status=row.get("status") or "unknown", + is_terminal=bool(row.get("is_terminal")), + realized_pnl_quote=net_pnl_quote if settled else 0.0, + unrealized_pnl_quote=0.0 if settled else net_pnl_quote, + global_pnl_quote=net_pnl_quote, + global_pnl_pct=_as_float(row.get("net_pnl_pct")), + # filled_amount_quote IS the volume traded, on every executor type including LP. + # There is no second volume column to reach for -- see + # test_executor_volume_is_the_filled_amount.py. + volume_quote=_as_float(row.get("filled_amount_quote")), + cum_fees_quote=_as_float(row.get("cum_fees_quote")), + controller_id=row.get("controller_id"), + executor_id=row.get("executor_id"), + executor_type=row.get("executor_type"), + account_name=row.get("account_name"), + connector_name=row.get("connector_name"), + trading_pair=row.get("trading_pair"), + close_type=row.get("close_type"), + ) diff --git a/pyproject.toml b/pyproject.toml index 4ad399a1..2c99deac 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -15,6 +15,11 @@ use_parentheses = true ensure_newline_before_comments = true combine_as_imports = true +[tool.pytest.ini_options] +testpaths = ["test"] +asyncio_mode = "auto" +asyncio_default_fixture_loop_scope = "function" + [tool.pre-commit] repos = [ { repo = "https://github.com/pre-commit/pre-commit-hooks", rev = "v3.4.0", hooks = [{ id = "check-yaml" }, { id = "end-of-file-fixer" }] }, diff --git a/routers/accounts.py b/routers/accounts.py index 3a80ad88..09881776 100644 --- a/routers/accounts.py +++ b/routers/accounts.py @@ -3,7 +3,7 @@ from fastapi import APIRouter, Depends, HTTPException from starlette import status -from deps import get_accounts_service +from deps import get_accounts_service, require_gateway_online from models import GatewayWalletCredential, SetDefaultWalletRequest from services.accounts_service import AccountsService, validate_safe_name @@ -93,7 +93,11 @@ async def delete_account(account_name: str, accounts_service: AccountsService = @router.post("/delete-credential/{account_name}/{connector_name}") -async def delete_credential(account_name: str, connector_name: str, accounts_service: AccountsService = Depends(get_accounts_service)): +async def delete_credential( + account_name: str, + connector_name: str, + accounts_service: AccountsService = Depends(get_accounts_service), +): """ Delete a specific connector credential for an account. @@ -115,7 +119,12 @@ async def delete_credential(account_name: str, connector_name: str, accounts_ser @router.post("/add-credential/{account_name}/{connector_name}", status_code=status.HTTP_201_CREATED) -async def add_credential(account_name: str, connector_name: str, credentials: Dict, accounts_service: AccountsService = Depends(get_accounts_service)): +async def add_credential( + account_name: str, + connector_name: str, + credentials: Dict, + accounts_service: AccountsService = Depends(get_accounts_service), +): """ Add or update connector credentials (API keys) for a specific account and connector. @@ -195,7 +204,7 @@ async def add_gateway_wallet( raise HTTPException(status_code=500, detail=str(e)) -@router.post("/gateway/wallet/set-default") +@router.post("/gateway/wallet/set-default", dependencies=[Depends(require_gateway_online)]) async def set_default_gateway_wallet( request: SetDefaultWalletRequest, accounts_service: AccountsService = Depends(get_accounts_service) @@ -219,9 +228,6 @@ async def set_default_gateway_wallet( } """ try: - if not await accounts_service.gateway_client.ping(): - raise HTTPException(status_code=503, detail="Gateway service is not available") - result = await accounts_service.gateway_client.set_default_wallet( chain=request.chain, address=request.address diff --git a/routers/archived_bots.py b/routers/archived_bots.py index 0b2bb93b..69455a2b 100644 --- a/routers/archived_bots.py +++ b/routers/archived_bots.py @@ -3,8 +3,8 @@ from fastapi import APIRouter, Depends, HTTPException, Query -from database import AsyncDatabaseManager, BotRunRepository -from deps import get_database_manager +from deps import get_bots_orchestrator +from services.bots_orchestrator import BotsOrchestrator from utils.file_system import fs_util from utils.hummingbot_database_reader import HummingbotDatabase @@ -41,7 +41,7 @@ async def list_databases(): @router.delete("/{db_path:path}") async def delete_archived_bot( db_path: str, - db_manager: AsyncDatabaseManager = Depends(get_database_manager) + bots_orchestrator: BotsOrchestrator = Depends(get_bots_orchestrator) ): """ Delete an archived bot and its entire directory. @@ -64,16 +64,9 @@ async def delete_archived_bot( except Exception as e: raise HTTPException(status_code=500, detail=f"Error deleting archived bot: {str(e)}") - # Best-effort: also clean matching BotRun records from PG - bot_runs_deleted = 0 - try: - async with db_manager.get_session_context() as session: - bot_run_repo = BotRunRepository(session) - bot_runs_deleted = await bot_run_repo.delete_bot_runs_by_bot_name(bot_name) - if bot_runs_deleted > 0: - logger.info(f"Deleted {bot_runs_deleted} bot run record(s) for '{bot_name}'") - except Exception as e: - logger.warning(f"Failed to clean bot run records for '{bot_name}': {e}") + # Best-effort: also clean matching BotRun records from PG. Bot-run persistence is + # the orchestrator's (ARCH-035), so the session and the swallow live there. + bot_runs_deleted = await bots_orchestrator.delete_bot_runs_for_bot(bot_name) return { "message": f"Archived bot '{bot_name}' deleted successfully", diff --git a/routers/backtesting.py b/routers/backtesting.py index ff8f2b5a..f22690bc 100644 --- a/routers/backtesting.py +++ b/routers/backtesting.py @@ -2,7 +2,7 @@ from deps import get_backtesting_service from models.backtesting import BacktestingConfig -from services.backtesting_service import BacktestingService +from services.backtesting_service import BacktestingService, BacktestTimeout router = APIRouter(tags=["Backtesting"], prefix="/backtesting") @@ -15,8 +15,16 @@ async def run_backtesting( """Run a backtest synchronously. Returns results directly (may timeout for long backtests).""" try: return await service.run_backtest_sync(backtesting_config.model_dump()) + except BacktestTimeout as e: + # A run abandoned at the wall-clock budget is a gateway timeout, not an engine + # error: the distinction is the only thing that tells a caller to retry smaller + # rather than to fix the config. + raise HTTPException(status_code=504, detail=str(e)) except Exception as e: - return {"error": str(e)} + # Answering 200 with {"error": ...} made every failure -- a bad controller, a + # dead worker, a terminated run -- indistinguishable from a finished backtest, + # and the api-client only raises on non-2xx, so nothing downstream ever noticed. + raise HTTPException(status_code=500, detail=str(e)) @router.post("/tasks") diff --git a/routers/bot_orchestration.py b/routers/bot_orchestration.py index ed76b98d..288d150f 100644 --- a/routers/bot_orchestration.py +++ b/routers/bot_orchestration.py @@ -417,44 +417,28 @@ async def stop_and_archive_bot( Returns immediately with a success message while the process continues in the background. """ try: - # Step 1: Normalize bot name and container name - # Container name is now the same as bot name (no prefix added) - actual_bot_name = bot_name - container_name = bot_name - - logging.info(f"Normalized bot_name: {actual_bot_name}, container_name: {container_name}") - - # Step 2: Validate bot exists in active bots + # The bot name is also the container name and the MQTT identity, so it is + # used verbatim everywhere below (no "hummingbot-" prefix is added anywhere). active_bots = list(bots_manager.active_bots.keys()) - # Check if bot exists in active bots (could be stored as either format) - bot_found = (actual_bot_name in active_bots) or (container_name in active_bots) - - if not bot_found: + if bot_name not in active_bots: return { "status": "error", "message": ( - f"Bot '{actual_bot_name}' not found in active bots. " + f"Bot '{bot_name}' not found in active bots. " f"Active bots: {active_bots}. Cannot perform graceful shutdown." ), "details": { - "input_name": bot_name, - "actual_bot_name": actual_bot_name, - "container_name": container_name, + "bot_name": bot_name, "active_bots": active_bots, "reason": "Bot must be actively managed via MQTT for graceful shutdown" } } - # Use the format that's actually stored in active bots - bot_name_for_orchestrator = container_name if container_name in active_bots else actual_bot_name - # Add the background task background_tasks.add_task( bots_manager.stop_and_archive_bot, - bot_name=actual_bot_name, - container_name=container_name, - bot_name_for_orchestrator=bot_name_for_orchestrator, + bot_name=bot_name, skip_order_cancellation=skip_order_cancellation, archive_locally=archive_locally, s3_bucket=s3_bucket, @@ -464,11 +448,9 @@ async def stop_and_archive_bot( return { "status": "success", - "message": f"Stop and archive process started for bot {actual_bot_name}", + "message": f"Stop and archive process started for bot {bot_name}", "details": { - "input_name": bot_name, - "actual_bot_name": actual_bot_name, - "container_name": container_name, + "bot_name": bot_name, "process": ( "The bot will be gracefully stopped, archived, and removed in the background. " "This process typically takes 20-30 seconds." diff --git a/routers/docker.py b/routers/docker.py index 7b0f8287..1a28c99b 100644 --- a/routers/docker.py +++ b/routers/docker.py @@ -1,42 +1,71 @@ import os -from fastapi import APIRouter, HTTPException, Depends +from fastapi import APIRouter, Depends, HTTPException +from deps import get_bot_archiver, get_docker_service from models import DockerImage -from utils.bot_archiver import BotArchiver from services.docker_service import DockerService -from deps import get_docker_service, get_bot_archiver +from utils.bot_archiver import BotArchiver router = APIRouter(tags=["Docker"], prefix="/docker") +# What a DockerService failure means over HTTP. The service reports the kind of failure +# rather than raising, because its internal callers (stop-and-archive's retry loop) are +# built to carry on past one; translating it is the HTTP layer's job. +_STATUS_BY_ERROR = { + "not_found": 404, + "docker_error": 502, +} + + +def _or_http_error(result): + """Turn a DockerService failure into a status code instead of a 200 with an error body. + + These routes used to return the service's error verbatim, so a container that does + not exist answered 200 with a raw docker-py string -- a caller checking the status + code, which is how a caller checks, read it as a container that had been stopped. + + Only a dict carrying `success: False` is a failure; a listing route's list and a + payload like {"images": [...]} pass through untouched. + """ + if isinstance(result, dict) and result.get("success") is False: + raise HTTPException( + status_code=_STATUS_BY_ERROR.get(result.get("error"), 502), + detail=result.get("message", "Docker operation failed"), + ) + return result + @router.get("/running") async def is_docker_running(docker_service: DockerService = Depends(get_docker_service)): """ Check if Docker daemon is running. - + Args: docker_service: Docker service dependency - + Returns: Dictionary indicating if Docker is running """ - return docker_service.is_docker_running() + return docker_service.is_docker_running() @router.get("/available-images/") async def available_images(image_name: str = None, docker_service: DockerService = Depends(get_docker_service)): """ Get available Docker images matching the specified name. - + Args: image_name: Name pattern to search for in image tags docker_service: Docker service dependency - + Returns: Dictionary with list of available image tags + + Raises: + HTTPException: 502 if the Docker daemon refused or could not be reached """ - available_images = docker_service.get_available_images() + available_images = _or_http_error(docker_service.get_available_images()) if image_name: return [tag for image in available_images["images"] for tag in image.tags if image_name in tag] return [tag for tag in available_images["images"]] @@ -46,79 +75,104 @@ async def available_images(image_name: str = None, docker_service: DockerService async def active_containers(name_filter: str = None, docker_service: DockerService = Depends(get_docker_service)): """ Get all currently active (running) Docker containers. - + Args: name_filter: Optional filter to match container names (case-insensitive) docker_service: Docker service dependency - + Returns: List of active container information + + Raises: + HTTPException: 502 if the Docker daemon refused or could not be reached """ - return docker_service.get_active_containers(name_filter) + return _or_http_error(docker_service.get_active_containers(name_filter)) @router.get("/exited-containers") async def exited_containers(name_filter: str = None, docker_service: DockerService = Depends(get_docker_service)): """ Get all exited (stopped) Docker containers. - + Args: name_filter: Optional filter to match container names (case-insensitive) docker_service: Docker service dependency - + Returns: List of exited container information + + Raises: + HTTPException: 502 if the Docker daemon refused or could not be reached """ - return docker_service.get_exited_containers(name_filter) + return _or_http_error(docker_service.get_exited_containers(name_filter)) @router.post("/clean-exited-containers") async def clean_exited_containers(docker_service: DockerService = Depends(get_docker_service)): """ Remove all exited Docker containers to free up space. - + Args: docker_service: Docker service dependency - + Returns: - Response from cleanup operation + {"success": true, "message": ...} once the exited containers are pruned + + Raises: + HTTPException: 502 if the Docker daemon refused or could not be reached """ - return docker_service.clean_exited_containers() + return _or_http_error(docker_service.clean_exited_containers()) @router.post("/remove-container/{container_name}") -async def remove_container(container_name: str, archive_locally: bool = True, s3_bucket: str = None, docker_service: DockerService = Depends(get_docker_service), bot_archiver: BotArchiver = Depends(get_bot_archiver)): +async def remove_container( + container_name: str, + archive_locally: bool = True, + s3_bucket: str = None, + docker_service: DockerService = Depends(get_docker_service), + bot_archiver: BotArchiver = Depends(get_bot_archiver), +): """ - Remove a Hummingbot container and optionally archive its bot data. - - NOTE: This endpoint only works with Hummingbot containers (names starting with 'hummingbot-') - as it archives bot-specific data from the bots/instances directory. - + Remove a bot container created by this API and archive its bot data. + + NOTE: This endpoint only works with containers this API manages. A bot container is named + after its instance verbatim, and owns the bots/instances/ directory that this + endpoint archives; an unrelated container on the host has no such directory and is refused. + Args: - container_name: Name of the Hummingbot container to remove + container_name: Name of the bot container to remove archive_locally: Whether to archive data locally (default: True) s3_bucket: S3 bucket name for cloud archiving (optional) docker_service: Docker service dependency bot_archiver: Bot archiver service dependency - + Returns: Response from container removal operation - + Raises: - HTTPException: 400 if container is not a Hummingbot container + HTTPException: 400 if the container is not a bot managed by this API HTTPException: 500 if archiving fails """ - # Validate that this is a Hummingbot container - if not container_name.startswith("hummingbot-"): + # Validate that this container belongs to a bot this API created, by the only marker that + # actually exists: its instance directory. Container names carry no prefix. + try: + instance_dir = DockerService.resolve_instance_dir(container_name) + except ValueError as e: + raise HTTPException(status_code=400, detail=str(e)) + if not os.path.isdir(instance_dir): raise HTTPException( - status_code=400, - detail=f"This endpoint only removes Hummingbot containers. Container '{container_name}' is not a Hummingbot container." + status_code=400, + detail=f"This endpoint only removes bot containers managed by this API. Container " + f"'{container_name}' has no instance directory at '{instance_dir}'." ) - + # Remove the container + # Deliberately not mapped to a status code the way stop/start are: this route also + # archives, and a container that is already gone is a bot whose data still needs + # archiving -- a removal that finds nothing to remove has reached the end state the + # caller asked for. The response carries `success` for a caller that wants to know + # which of the two happened. response = docker_service.remove_container(container_name) - # Form the instance directory path correctly - instance_dir = os.path.join('bots', 'instances', container_name) try: # Archive the data if archive_locally: @@ -135,30 +189,38 @@ async def remove_container(container_name: str, archive_locally: bool = True, s3 async def stop_container(container_name: str, docker_service: DockerService = Depends(get_docker_service)): """ Stop a running Docker container. - + Args: container_name: Name of the container to stop docker_service: Docker service dependency - + Returns: - Response from container stop operation + {"success": true, "message": ...} once the container is stopped + + Raises: + HTTPException: 404 if no container by that name exists + HTTPException: 502 if the Docker daemon refused or could not be reached """ - return docker_service.stop_container(container_name) + return _or_http_error(docker_service.stop_container(container_name)) @router.post("/start-container/{container_name}") async def start_container(container_name: str, docker_service: DockerService = Depends(get_docker_service)): """ Start a stopped Docker container. - + Args: container_name: Name of the container to start docker_service: Docker service dependency - + Returns: - Response from container start operation + {"success": true, "message": ...} once the container is started + + Raises: + HTTPException: 404 if no container by that name exists + HTTPException: 502 if the Docker daemon refused or could not be reached """ - return docker_service.start_container(container_name) + return _or_http_error(docker_service.start_container(container_name)) @router.post("/pull-image/") @@ -166,11 +228,11 @@ async def pull_image(image: DockerImage, docker_service: DockerService = Depends """ Initiate Docker image pull as background task. Returns immediately with task status for monitoring. - + Args: image: DockerImage object containing the image name to pull docker_service: Docker service dependency - + Returns: Status of the pull operation initiation """ @@ -182,10 +244,10 @@ async def pull_image(image: DockerImage, docker_service: DockerService = Depends async def get_pull_status(docker_service: DockerService = Depends(get_docker_service)): """ Get status of all pull operations. - + Args: docker_service: Docker service dependency - + Returns: Dictionary with all pull operations and their statuses """ diff --git a/routers/gateway.py b/routers/gateway.py index 1ee9ea3a..945a60d9 100644 --- a/routers/gateway.py +++ b/routers/gateway.py @@ -3,7 +3,7 @@ from fastapi import APIRouter, Depends, HTTPException, Query -from deps import get_accounts_service, get_gateway_service +from deps import get_accounts_service, get_gateway_service, require_gateway_online from models import AddPoolRequest, AddTokenRequest, GatewayConfig, GatewayStatus, UpdateApiKeysRequest from services.accounts_service import AccountsService from services.gateway_client import GatewayError, check_gateway_error @@ -154,7 +154,7 @@ async def get_gateway_logs( # Connectors # ============================================ -@router.get("/connectors") +@router.get("/connectors", dependencies=[Depends(require_gateway_online)]) async def list_connectors(accounts_service: AccountsService = Depends(get_accounts_service)) -> Dict: """ List all available DEX connectors with their configurations. @@ -163,9 +163,6 @@ async def list_connectors(accounts_service: AccountsService = Depends(get_accoun All fields normalized to snake_case. """ try: - if not await accounts_service.gateway_client.ping(): - raise HTTPException(status_code=503, detail="Gateway service is not available") - result = check_gateway_error(await accounts_service.gateway_client.get_connectors()) return normalize_gateway_response(result) @@ -177,7 +174,7 @@ async def list_connectors(accounts_service: AccountsService = Depends(get_accoun raise HTTPException(status_code=500, detail=f"Error listing connectors: {str(e)}") -@router.get("/connectors/{connector_name}") +@router.get("/connectors/{connector_name}", dependencies=[Depends(require_gateway_online)]) async def get_connector_config( connector_name: str, accounts_service: AccountsService = Depends(get_accounts_service) @@ -189,9 +186,6 @@ async def get_connector_config( connector_name: Connector name (e.g., 'meteora', 'raydium') """ try: - if not await accounts_service.gateway_client.ping(): - raise HTTPException(status_code=503, detail="Gateway service is not available") - result = check_gateway_error(await accounts_service.gateway_client.get_config(connector_name)) return normalize_gateway_response(result) @@ -203,7 +197,7 @@ async def get_connector_config( raise HTTPException(status_code=500, detail=f"Error getting connector config: {str(e)}") -@router.post("/connectors/{connector_name}") +@router.post("/connectors/{connector_name}", dependencies=[Depends(require_gateway_online)]) async def update_connector_config( connector_name: str, config_updates: Dict, @@ -219,9 +213,6 @@ async def update_connector_config( or camelCase (e.g., {"slippagePct": 0.5}) """ try: - if not await accounts_service.gateway_client.ping(): - raise HTTPException(status_code=503, detail="Gateway service is not available") - results = [] for path, value in config_updates.items(): # Convert snake_case to camelCase if needed @@ -251,7 +242,7 @@ async def update_connector_config( # API Keys # ============================================ -@router.get("/apiKeys") +@router.get("/apiKeys", dependencies=[Depends(require_gateway_online)]) async def get_api_keys(accounts_service: AccountsService = Depends(get_accounts_service)) -> Dict: """ Get all configured API keys from Gateway. @@ -266,9 +257,6 @@ async def get_api_keys(accounts_service: AccountsService = Depends(get_accounts_ } """ try: - if not await accounts_service.gateway_client.ping(): - raise HTTPException(status_code=503, detail="Gateway service is not available") - return check_gateway_error(await accounts_service.gateway_client.get_api_keys()) except HTTPException: @@ -279,7 +267,7 @@ async def get_api_keys(accounts_service: AccountsService = Depends(get_accounts_ raise HTTPException(status_code=500, detail=f"Error getting API keys: {str(e)}") -@router.post("/apiKeys") +@router.post("/apiKeys", dependencies=[Depends(require_gateway_online)]) async def update_api_keys( request: UpdateApiKeysRequest, accounts_service: AccountsService = Depends(get_accounts_service) @@ -301,9 +289,6 @@ async def update_api_keys( Note: After updating API keys, restart Gateway for changes to take effect. """ try: - if not await accounts_service.gateway_client.ping(): - raise HTTPException(status_code=503, detail="Gateway service is not available") - results = await accounts_service.gateway_client.update_api_keys(request.api_keys) # A None result means the request never reached Gateway (connection @@ -335,7 +320,7 @@ async def update_api_keys( # Chains (Networks) and Tokens # ============================================ -@router.get("/chains") +@router.get("/chains", dependencies=[Depends(require_gateway_online)]) async def list_chains(accounts_service: AccountsService = Depends(get_accounts_service)) -> Dict: """ List all available blockchain chains and their networks. @@ -343,9 +328,6 @@ async def list_chains(accounts_service: AccountsService = Depends(get_accounts_s This also serves as the networks list endpoint. """ try: - if not await accounts_service.gateway_client.ping(): - raise HTTPException(status_code=503, detail="Gateway service is not available") - return check_gateway_error(await accounts_service.gateway_client.get_chains()) except HTTPException: @@ -360,7 +342,7 @@ async def list_chains(accounts_service: AccountsService = Depends(get_accounts_s # Pools (Legacy - use /networks/{network_id}/pools instead) # ============================================ -@router.get("/pools", deprecated=True) +@router.get("/pools", deprecated=True, dependencies=[Depends(require_gateway_online)]) async def list_pools_legacy( connector_name: str = Query(description="DEX connector (e.g., 'meteora', 'raydium')"), network: str = Query(description="Network (e.g., 'mainnet-beta')"), @@ -372,9 +354,6 @@ async def list_pools_legacy( List all liquidity pools for a connector and network. """ try: - if not await accounts_service.gateway_client.ping(): - raise HTTPException(status_code=503, detail="Gateway service is not available") - # Determine chain from connector (legacy behavior). Solana routers # (jupiter/dflow/okx/titan) must map to solana too, or this returns the # (empty) ethereum pool list for them. @@ -404,7 +383,7 @@ async def list_pools_legacy( # Networks (Primary Endpoints) # ============================================ -@router.get("/networks") +@router.get("/networks", dependencies=[Depends(require_gateway_online)]) async def list_networks(accounts_service: AccountsService = Depends(get_accounts_service)) -> Dict: """ List all available networks across all chains. @@ -413,9 +392,6 @@ async def list_networks(accounts_service: AccountsService = Depends(get_accounts This is the primary interface for network discovery. """ try: - if not await accounts_service.gateway_client.ping(): - raise HTTPException(status_code=503, detail="Gateway service is not available") - chains_result = check_gateway_error(await accounts_service.gateway_client.get_chains()) # Flatten chain-network combinations into network IDs @@ -445,7 +421,7 @@ async def list_networks(accounts_service: AccountsService = Depends(get_accounts raise HTTPException(status_code=500, detail=f"Error listing networks: {str(e)}") -@router.get("/networks/{network_id}") +@router.get("/networks/{network_id}", dependencies=[Depends(require_gateway_online)]) async def get_network_config( network_id: str, accounts_service: AccountsService = Depends(get_accounts_service) @@ -459,9 +435,6 @@ async def get_network_config( Example: GET /gateway/networks/solana-mainnet-beta """ try: - if not await accounts_service.gateway_client.ping(): - raise HTTPException(status_code=503, detail="Gateway service is not available") - result = check_gateway_error(await accounts_service.gateway_client.get_config(network_id)) return normalize_gateway_response(result) @@ -473,7 +446,7 @@ async def get_network_config( raise HTTPException(status_code=500, detail=f"Error getting network config: {str(e)}") -@router.post("/networks/{network_id}") +@router.post("/networks/{network_id}", dependencies=[Depends(require_gateway_online)]) async def update_network_config( network_id: str, config_updates: Dict, @@ -491,9 +464,6 @@ async def update_network_config( Example: POST /gateway/networks/solana-mainnet-beta """ try: - if not await accounts_service.gateway_client.ping(): - raise HTTPException(status_code=503, detail="Gateway service is not available") - results = [] for path, value in config_updates.items(): # Convert snake_case to camelCase if needed @@ -519,7 +489,7 @@ async def update_network_config( raise HTTPException(status_code=500, detail=f"Error updating network config: {str(e)}") -@router.get("/networks/{network_id}/tokens") +@router.get("/networks/{network_id}/tokens", dependencies=[Depends(require_gateway_online)]) async def get_network_tokens( network_id: str, search: Optional[str] = Query(default=None), @@ -535,9 +505,6 @@ async def get_network_tokens( Example: GET /gateway/networks/solana-mainnet-beta/tokens?search=USDC """ try: - if not await accounts_service.gateway_client.ping(): - raise HTTPException(status_code=503, detail="Gateway service is not available") - # Parse network_id into chain and network parts = network_id.split('-', 1) if len(parts) != 2: @@ -572,7 +539,7 @@ async def get_network_tokens( raise HTTPException(status_code=500, detail=f"Error getting network tokens: {str(e)}") -@router.post("/networks/{network_id}/tokens") +@router.post("/networks/{network_id}/tokens", dependencies=[Depends(require_gateway_online)]) async def add_network_token( network_id: str, token_request: AddTokenRequest, @@ -597,9 +564,6 @@ async def add_network_token( the token is live as soon as this returns. """ try: - if not await accounts_service.gateway_client.ping(): - raise HTTPException(status_code=503, detail="Gateway service is not available") - # Parse network_id into chain and network parts = network_id.split('-', 1) if len(parts) != 2: @@ -637,7 +601,7 @@ async def add_network_token( raise HTTPException(status_code=500, detail=f"Error adding token: {str(e)}") -@router.post("/networks/{network_id}/tokens/save/{token_address}") +@router.post("/networks/{network_id}/tokens/save/{token_address}", dependencies=[Depends(require_gateway_online)]) async def save_network_token( network_id: str, token_address: str, @@ -656,9 +620,6 @@ async def save_network_token( Example: POST /gateway/networks/solana-mainnet-beta/tokens/save/9QFfgxdSqH5zT7j6rZb1y6SZhw2aFtcQu2r6BuYpump """ try: - if not await accounts_service.gateway_client.ping(): - raise HTTPException(status_code=503, detail="Gateway service is not available") - # Parse network_id into chain and network parts = network_id.split('-', 1) if len(parts) != 2: @@ -693,7 +654,7 @@ async def save_network_token( raise HTTPException(status_code=500, detail=f"Error saving token: {str(e)}") -@router.delete("/networks/{network_id}/tokens/{token_address}") +@router.delete("/networks/{network_id}/tokens/{token_address}", dependencies=[Depends(require_gateway_online)]) async def delete_network_token( network_id: str, token_address: str, @@ -712,9 +673,6 @@ async def delete_network_token( the deletion is live as soon as this returns. """ try: - if not await accounts_service.gateway_client.ping(): - raise HTTPException(status_code=503, detail="Gateway service is not available") - # Parse network_id into chain and network parts = network_id.split('-', 1) if len(parts) != 2: @@ -747,7 +705,7 @@ async def delete_network_token( # Network Pools # ============================================ -@router.get("/networks/{network_id}/pools") +@router.get("/networks/{network_id}/pools", dependencies=[Depends(require_gateway_online)]) async def get_network_pools( network_id: str, connector: Optional[str] = Query(default=None, description="Filter by connector (e.g., 'raydium', 'meteora')"), @@ -767,9 +725,6 @@ async def get_network_pools( Example: GET /gateway/networks/solana-mainnet-beta/pools?connector=raydium&type=clmm """ try: - if not await accounts_service.gateway_client.ping(): - raise HTTPException(status_code=503, detail="Gateway service is not available") - # Parse network_id into chain and network parts = network_id.split('-', 1) if len(parts) != 2: @@ -801,7 +756,7 @@ async def get_network_pools( raise HTTPException(status_code=500, detail=f"Error getting network pools: {str(e)}") -@router.post("/networks/{network_id}/pools") +@router.post("/networks/{network_id}/pools", dependencies=[Depends(require_gateway_online)]) async def add_network_pool( network_id: str, pool_request: AddPoolRequest, @@ -830,9 +785,6 @@ async def add_network_pool( pool is listed and priced as soon as this returns. """ try: - if not await accounts_service.gateway_client.ping(): - raise HTTPException(status_code=503, detail="Gateway service is not available") - # Parse network_id into chain and network parts = network_id.split('-', 1) if len(parts) != 2: @@ -879,7 +831,7 @@ async def add_network_pool( raise HTTPException(status_code=500, detail=f"Error adding pool: {str(e)}") -@router.post("/networks/{network_id}/pools/save/{pool_address}") +@router.post("/networks/{network_id}/pools/save/{pool_address}", dependencies=[Depends(require_gateway_online)]) async def save_network_pool( network_id: str, pool_address: str, @@ -906,9 +858,6 @@ async def save_network_pool( Example: POST /gateway/networks/solana-mainnet-beta/pools/save/2sf5NYcY...?connector=meteora&type=clmm """ try: - if not await accounts_service.gateway_client.ping(): - raise HTTPException(status_code=503, detail="Gateway service is not available") - if (connector is None) != (type is None): raise HTTPException( status_code=400, @@ -947,7 +896,7 @@ async def save_network_pool( raise HTTPException(status_code=500, detail=f"Error saving pool: {str(e)}") -@router.delete("/networks/{network_id}/pools/{pool_address}") +@router.delete("/networks/{network_id}/pools/{pool_address}", dependencies=[Depends(require_gateway_online)]) async def delete_network_pool( network_id: str, pool_address: str, @@ -966,9 +915,6 @@ async def delete_network_pool( deletion is live as soon as this returns. """ try: - if not await accounts_service.gateway_client.ping(): - raise HTTPException(status_code=503, detail="Gateway service is not available") - # Parse network_id into chain and network parts = network_id.split('-', 1) if len(parts) != 2: diff --git a/routers/gateway_amm.py b/routers/gateway_amm.py index 8cbb64e4..3c6639e2 100644 --- a/routers/gateway_amm.py +++ b/routers/gateway_amm.py @@ -20,15 +20,12 @@ position_address, and Gateway rejects positions-owned for them with a 400, surfaced as-is. """ import logging -from decimal import Decimal from typing import Any, Dict, List, Optional from fastapi import APIRouter, Depends, HTTPException -from database import AsyncDatabaseManager -from database.repositories import GatewayAMMRepository from database.repositories.gateway_amm_repository import has_nft_positions -from deps import get_accounts_service, get_database_manager +from deps import get_accounts_service, get_gateway_amm_service, require_gateway_online from models import ( AMMAddLiquidityRequest, AMMCreatePoolRequest, @@ -43,12 +40,14 @@ ) from routers.gateway_extras import ( ExtraParamsSpec, + get_transaction_hash_from_response, get_transaction_status_from_response, transaction_id_from_error, validate_extra_params, ) from services.accounts_service import AccountsService -from services.gateway_client import GatewayError, check_gateway_error, get_native_gas_token +from services.gateway_amm_service import GatewayAMMService +from services.gateway_client import GatewayError, check_gateway_error logger = logging.getLogger(__name__) @@ -63,11 +62,6 @@ } -async def _require_gateway(accounts_service: AccountsService) -> None: - if not await accounts_service.gateway_client.ping(): - raise HTTPException(status_code=503, detail="Gateway service is not available") - - async def _resolve_wallet(accounts_service: AccountsService, network: str, wallet_address) -> str: chain, _ = accounts_service.gateway_client.parse_network_id(network) return await accounts_service.gateway_client.get_wallet_address_or_default( @@ -99,7 +93,7 @@ async def _read_pool( async def _book_position_add( accounts_service: AccountsService, - db_manager: AsyncDatabaseManager, + amm_service: GatewayAMMService, request: AMMAddLiquidityRequest, wallet_address: str, position_address: str, @@ -107,60 +101,29 @@ async def _book_position_add( pool_info: Dict[str, Any], price: Optional[float], ) -> None: - """Create or top up the DAMM v2 position row for a confirmed add.""" - base_added = data.get("baseTokenAmountAdded") or 0 - quote_added = data.get("quoteTokenAmountAdded") or 0 - # Rent is locked, not spent: the chain returns it when the position account closes, - # and Gateway reports it separately for exactly that reason. Recorded here so the - # close can be checked against it — a refund smaller than what was locked means an - # account was left behind. Present only when this add opened the position; adding to - # one that already exists locks no further rent. - position_rent = data.get("positionRent") + """Pull the add's amounts out of Gateway's response and hand them to the service. - try: - async with db_manager.get_session_context() as session: - repo = GatewayAMMRepository(session) - existing = await repo.get_position_by_address(position_address) - if existing: - await repo.add_to_position_amounts( - position_address=position_address, - base_delta=Decimal(str(base_added)), - quote_delta=Decimal(str(quote_added)), - entry_price=Decimal(str(price)) if price else None, - ) - if existing.status == "CLOSED": - existing.status = "OPEN" - existing.closed_at = None - else: - chain, network_name = request.network.split("-", 1) - base_symbol = await accounts_service.gateway_client.resolve_token_symbol( - chain, network_name, pool_info.get("baseTokenAddress", "")) - quote_symbol = await accounts_service.gateway_client.resolve_token_symbol( - chain, network_name, pool_info.get("quoteTokenAddress", "")) - await repo.create_position({ - "position_address": position_address, - "pool_address": request.pool_address, - "connector": request.connector.split("/")[0], - "network": request.network, - "wallet_address": wallet_address, - "base_token": base_symbol, - "quote_token": quote_symbol, - "trading_pair": f"{base_symbol}-{quote_symbol}", - "initial_base_token_amount": base_added, - "initial_quote_token_amount": quote_added, - "base_token_amount": base_added, - "quote_token_amount": quote_added, - "position_rent": Decimal(str(position_rent)) if position_rent is not None else None, - "entry_price": price, - "current_price": price, - }) - logger.info(f"Booked AMM position {position_address}: +{base_added} base, +{quote_added} quote") - except Exception as db_error: - logger.error(f"Error booking AMM position {position_address}: {db_error}", exc_info=True) + Reading the response is the router's job; what row the position becomes is the + service's. ``positionRent`` is present only when this add opened the position. + """ + await amm_service.record_position_add( + gateway_client=accounts_service.gateway_client, + position_address=position_address, + pool_address=request.pool_address, + connector=request.connector, + network=request.network, + wallet_address=wallet_address, + base_amount_added=data.get("baseTokenAmountAdded"), + quote_amount_added=data.get("quoteTokenAmountAdded"), + position_rent=data.get("positionRent"), + price=price, + base_token_address=pool_info.get("baseTokenAddress", ""), + quote_token_address=pool_info.get("quoteTokenAddress", ""), + ) async def _record_event( - db_manager: AsyncDatabaseManager, + amm_service: GatewayAMMService, result: Dict[str, Any], *, event_type: str, @@ -173,38 +136,33 @@ async def _record_event( quote_amount_key: str, price: Optional[float] = None, ) -> str: - """Persist one AMM write and return its status in hapi's vocabulary. + """Hand one AMM write to the service and return its status in hapi's vocabulary. - Amounts come from Gateway's `data`, present only once it confirmed the tx; a - submitted-not-confirmed write records the status with null amounts rather than - inventing figures. Recording never fails the operation — the liquidity has already - moved by the time we get here, so a database problem must not surface as a failed - write to the caller. + Gateway's response shape is parsed here — the amount keys differ per event type and + ``data`` is present only once it confirmed the tx, so a submitted-not-confirmed write + is recorded with null amounts rather than invented figures. Recording never fails the + operation; the service owns that policy. """ tx_status = get_transaction_status_from_response(result) data = result.get("data") or {} - chain, _ = network.split("-", 1) if "-" in network else (network, "") - try: - async with db_manager.get_session_context() as session: - await GatewayAMMRepository(session).create_event({ - "transaction_hash": result.get("signature") or result.get("txHash") or "", - "connector": connector, - "network": network, - "wallet_address": wallet_address, - "pool_address": pool_address, - "position_address": position_address, - "event_type": event_type, - "base_token_amount": data.get(base_amount_key), - "quote_token_amount": data.get(quote_amount_key), - "price": price, - "gas_fee": data.get("fee"), - "gas_token": get_native_gas_token(chain) if data.get("fee") is not None else None, - "status": tx_status, - }) - logger.info(f"Recorded AMM {event_type}: {result.get('signature')} (status: {tx_status})") - except Exception as db_error: - logger.error(f"Error recording AMM {event_type} event: {db_error}", exc_info=True) + await amm_service.record_event( + # "" rather than None: the column is not nullable, and a write whose id Gateway + # withheld is still worth recording — unlike the routes above, this one never + # fails the operation over a missing hash. + transaction_hash=get_transaction_hash_from_response(result) or "", + event_type=event_type, + connector=connector, + network=network, + wallet_address=wallet_address, + pool_address=pool_address, + position_address=position_address, + base_token_amount=data.get(base_amount_key), + quote_token_amount=data.get(quote_amount_key), + price=price, + gas_fee=data.get("fee"), + tx_status=tx_status, + ) return tx_status @@ -213,7 +171,7 @@ async def _record_event( async def _record_failed_event( - db_manager: AsyncDatabaseManager, + amm_service: GatewayAMMService, error: Exception, *, event_type: str, @@ -223,45 +181,31 @@ async def _record_failed_event( pool_address: str, position_address: Optional[str] = None, ) -> None: - """Record a write that reached the chain and reverted, before the error is re-raised. - - `_record_event` above only runs when Gateway *returns*. A transaction that landed and - reverted does not return: Gateway raises, the client turns it into a GatewayError, and - control skips the whole recording block. That is why every row in both event tables - read CONFIRMED with no error_message — not because nothing had ever failed, but - because a failure could not be written. + """Pull the transaction id out of a Gateway failure and hand it to the service. - Only failures carrying a transaction id are recorded: a pre-flight simulation failure - never got one and cost nothing, while a landed revert has one and paid gas. Recording - never masks the original failure. + A write that landed on-chain and reverted does not return: Gateway raises, and the + transaction id survives only inside the error message. Parsing it is the router's job + (one parser, shared with the success paths); deciding what row it becomes is the + service's. """ - transaction_hash = transaction_id_from_error(error) - if not transaction_hash: - return - - chain, _ = network.split("-", 1) if "-" in network else (network, "") - try: - async with db_manager.get_session_context() as session: - await GatewayAMMRepository(session).create_event({ - "transaction_hash": transaction_hash, - "connector": connector, - "network": network, - "wallet_address": wallet_address, - "pool_address": pool_address, - "position_address": position_address, - "event_type": event_type, - "status": "FAILED", - "error_message": str(error), - }) - logger.error( - f"AMM {event_type} {transaction_hash} landed on-chain and FAILED on {connector}/" - f"{network}; recorded. {error}" - ) - except Exception as db_error: - logger.error(f"Error recording failed AMM {event_type}: {db_error}", exc_info=True) + await amm_service.record_failed_event( + transaction_hash=transaction_id_from_error(error), + error=error, + event_type=event_type, + connector=connector, + network=network, + wallet_address=wallet_address, + pool_address=pool_address, + position_address=position_address, + ) -@router.get("/amm/pool-info", response_model=AMMPoolInfoResponse, response_model_by_alias=False) +@router.get( + "/amm/pool-info", + response_model=AMMPoolInfoResponse, + response_model_by_alias=False, + dependencies=[Depends(require_gateway_online)], +) async def get_amm_pool_info( connector: str, network: str, @@ -270,7 +214,6 @@ async def get_amm_pool_info( ): """Get AMM pool information (reserves, price, base fee) by pool address.""" try: - await _require_gateway(accounts_service) result = check_gateway_error(await accounts_service.gateway_client.amm_pool_info( connector=connector, chain_network=network, pool_address=pool_address, )) @@ -286,7 +229,12 @@ async def get_amm_pool_info( raise HTTPException(status_code=500, detail=f"Error getting AMM pool info: {str(e)}") -@router.get("/amm/position-info", response_model=AMMPositionInfoResponse, response_model_by_alias=False) +@router.get( + "/amm/position-info", + response_model=AMMPositionInfoResponse, + response_model_by_alias=False, + dependencies=[Depends(require_gateway_online)], +) async def get_amm_position_info( connector: str, network: str, @@ -296,7 +244,6 @@ async def get_amm_position_info( ): """Get a wallet's aggregate liquidity in an AMM pool plus a per-position breakdown (DAMM v2).""" try: - await _require_gateway(accounts_service) wallet_address = await _resolve_wallet(accounts_service, network, wallet_address) result = check_gateway_error(await accounts_service.gateway_client.amm_position_info( connector=connector, chain_network=network, pool_address=pool_address, @@ -314,7 +261,12 @@ async def get_amm_position_info( raise HTTPException(status_code=500, detail=f"Error getting AMM position info: {str(e)}") -@router.post("/amm/positions-owned", response_model=List[AMMPositionInfoResponse], response_model_by_alias=False) +@router.post( + "/amm/positions-owned", + response_model=List[AMMPositionInfoResponse], + response_model_by_alias=False, + dependencies=[Depends(require_gateway_online)], +) async def get_amm_positions_owned( request: AMMPositionsOwnedRequest, accounts_service: AccountsService = Depends(get_accounts_service), @@ -326,7 +278,6 @@ async def get_amm_positions_owned( them with a 400, surfaced here unchanged. Use position-info with a specific pool address instead. """ try: - await _require_gateway(accounts_service) wallet_address = await _resolve_wallet(accounts_service, request.network, request.wallet_address) result = check_gateway_error(await accounts_service.gateway_client.amm_positions_owned( connector=request.connector, chain_network=request.network, wallet_address=wallet_address, @@ -344,14 +295,18 @@ async def get_amm_positions_owned( raise HTTPException(status_code=500, detail=f"Error getting AMM positions owned: {str(e)}") -@router.post("/amm/quote-liquidity", response_model=AMMQuoteLiquidityResponse, response_model_by_alias=False) +@router.post( + "/amm/quote-liquidity", + response_model=AMMQuoteLiquidityResponse, + response_model_by_alias=False, + dependencies=[Depends(require_gateway_online)], +) async def quote_amm_liquidity( request: AMMQuoteLiquidityRequest, accounts_service: AccountsService = Depends(get_accounts_service), ): """Quote a two-sided liquidity deposit.""" try: - await _require_gateway(accounts_service) result = check_gateway_error(await accounts_service.gateway_client.amm_quote_liquidity( connector=request.connector, chain_network=request.network, pool_address=request.pool_address, base_token_amount=float(request.base_token_amount), quote_token_amount=float(request.quote_token_amount), @@ -371,11 +326,11 @@ async def quote_amm_liquidity( # ----------------------------- Writes ----------------------------- -@router.post("/amm/add-liquidity", response_model=AMMTransactionResponse) +@router.post("/amm/add-liquidity", response_model=AMMTransactionResponse, dependencies=[Depends(require_gateway_online)]) async def add_amm_liquidity( request: AMMAddLiquidityRequest, accounts_service: AccountsService = Depends(get_accounts_service), - db_manager: AsyncDatabaseManager = Depends(get_database_manager), + amm_service: GatewayAMMService = Depends(get_gateway_amm_service), ): """ Add two-sided liquidity to an AMM pool. @@ -387,7 +342,6 @@ async def add_amm_liquidity( # over the top of Gateway's own error when the wallet lookup is what failed. wallet_address = "" try: - await _require_gateway(accounts_service) wallet_address = await _resolve_wallet(accounts_service, request.network, request.wallet_address) result = check_gateway_error(await accounts_service.gateway_client.amm_add_liquidity( connector=request.connector, chain_network=request.network, wallet_address=wallet_address, @@ -417,11 +371,11 @@ async def add_amm_liquidity( position_address = data.get("positionAddress") if position_address: await _book_position_add( - accounts_service, db_manager, request, wallet_address, position_address, + accounts_service, amm_service, request, wallet_address, position_address, data, pool_info, price) tx_status = await _record_event( - db_manager, result, + amm_service, result, event_type="ADD_LIQUIDITY", connector=request.connector, network=request.network, wallet_address=wallet_address, pool_address=request.pool_address, position_address=position_address, @@ -433,7 +387,7 @@ async def add_amm_liquidity( raise except GatewayError as e: await _record_failed_event( - db_manager, e, event_type="ADD_LIQUIDITY", connector=request.connector, + amm_service, e, event_type="ADD_LIQUIDITY", connector=request.connector, network=request.network, wallet_address=wallet_address, pool_address=request.pool_address, position_address=request.position_address, ) @@ -445,11 +399,11 @@ async def add_amm_liquidity( raise HTTPException(status_code=500, detail=f"Error adding AMM liquidity: {str(e)}") -@router.post("/amm/remove-liquidity", response_model=AMMTransactionResponse) +@router.post("/amm/remove-liquidity", response_model=AMMTransactionResponse, dependencies=[Depends(require_gateway_online)]) async def remove_amm_liquidity( request: AMMRemoveLiquidityRequest, accounts_service: AccountsService = Depends(get_accounts_service), - db_manager: AsyncDatabaseManager = Depends(get_database_manager), + amm_service: GatewayAMMService = Depends(get_gateway_amm_service), ): """ Remove liquidity from an AMM pool. @@ -461,7 +415,6 @@ async def remove_amm_liquidity( # over the top of Gateway's own error when the wallet lookup is what failed. wallet_address = "" try: - await _require_gateway(accounts_service) wallet_address = await _resolve_wallet(accounts_service, request.network, request.wallet_address) result = check_gateway_error(await accounts_service.gateway_client.amm_remove_liquidity( connector=request.connector, chain_network=request.network, wallet_address=wallet_address, @@ -476,32 +429,16 @@ async def remove_amm_liquidity( if (get_transaction_status_from_response(result) == "CONFIRMED" and has_nft_positions(request.connector) and request.position_address): - try: - async with db_manager.get_session_context() as session: - repo = GatewayAMMRepository(session) - position = await repo.subtract_from_position_amounts( - position_address=request.position_address, - base_delta=Decimal(str(data.get("baseTokenAmountRemoved") or 0)), - quote_delta=Decimal(str(data.get("quoteTokenAmountRemoved") or 0)), - ) - # A 100% remove is the close: Gateway closes the position account in - # the same transaction, which is what returns its rent. There is no - # separate close route, and positionRentRefunded arrives only on this - # path — a partial removal leaves the account open and refunds - # nothing, so its absence there is a fact rather than a gap. - if position and float(request.percentage_to_remove) >= 100: - rent_refunded = data.get("positionRentRefunded") - await repo.close_position( - request.position_address, - position_rent_refunded=(Decimal(str(rent_refunded)) - if rent_refunded is not None else None), - ) - except Exception as db_error: - logger.error(f"Error booking AMM removal for {request.position_address}: " - f"{db_error}", exc_info=True) + await amm_service.record_position_remove( + position_address=request.position_address, + base_amount_removed=data.get("baseTokenAmountRemoved"), + quote_amount_removed=data.get("quoteTokenAmountRemoved"), + percentage_to_remove=float(request.percentage_to_remove), + position_rent_refunded=data.get("positionRentRefunded"), + ) tx_status = await _record_event( - db_manager, result, + amm_service, result, event_type="REMOVE_LIQUIDITY", connector=request.connector, network=request.network, wallet_address=wallet_address, pool_address=request.pool_address, position_address=request.position_address, @@ -513,7 +450,7 @@ async def remove_amm_liquidity( raise except GatewayError as e: await _record_failed_event( - db_manager, e, event_type="REMOVE_LIQUIDITY", connector=request.connector, + amm_service, e, event_type="REMOVE_LIQUIDITY", connector=request.connector, network=request.network, wallet_address=wallet_address, pool_address=request.pool_address, position_address=request.position_address, ) @@ -525,11 +462,16 @@ async def remove_amm_liquidity( raise HTTPException(status_code=500, detail=f"Error removing AMM liquidity: {str(e)}") -@router.post("/amm/create-pool", response_model=AMMCreatePoolResponse, response_model_by_alias=False) +@router.post( + "/amm/create-pool", + response_model=AMMCreatePoolResponse, + response_model_by_alias=False, + dependencies=[Depends(require_gateway_online)], +) async def create_amm_pool( request: AMMCreatePoolRequest, accounts_service: AccountsService = Depends(get_accounts_service), - db_manager: AsyncDatabaseManager = Depends(get_database_manager), + amm_service: GatewayAMMService = Depends(get_gateway_amm_service), ): """ Create and seed a new AMM pool. @@ -550,7 +492,6 @@ async def create_amm_pool( "(DAMM v2 pools are created against a config account)." ) - await _require_gateway(accounts_service) wallet_address = await _resolve_wallet(accounts_service, request.network, request.wallet_address) result = check_gateway_error(await accounts_service.gateway_client.amm_create_pool( connector=request.connector, chain_network=request.network, wallet_address=wallet_address, @@ -564,7 +505,7 @@ async def create_amm_pool( # The pool address only exists in the response, so it is read from there # rather than the request, which names tokens. tx_status = await _record_event( - db_manager, result, + amm_service, result, event_type="CREATE_POOL", connector=request.connector, network=request.network, wallet_address=wallet_address, pool_address=result.get("poolAddress") or result.get("pool_address") or "", @@ -595,7 +536,7 @@ async def search_amm_events( status: Optional[str] = None, limit: int = 50, offset: int = 0, - db_manager: AsyncDatabaseManager = Depends(get_database_manager), + amm_service: GatewayAMMService = Depends(get_gateway_amm_service), ): """ Search recorded AMM liquidity writes, newest first. @@ -605,19 +546,11 @@ async def search_amm_events( /gateway/amm/position-info, which is the only authority on them. """ try: - async with db_manager.get_session_context() as session: - repo = GatewayAMMRepository(session) - events = await repo.search_events( - connector=connector, network=network, wallet_address=wallet_address, - pool_address=pool_address, event_type=event_type, status=status, - limit=min(limit, 1000), offset=offset, - ) - return { - "data": [repo.event_to_dict(event) for event in events], - "total_count": len(events), - "limit": limit, - "offset": offset, - } + return await amm_service.search_events( + connector=connector, network=network, wallet_address=wallet_address, + pool_address=pool_address, event_type=event_type, status=status, + limit=limit, offset=offset, + ) except Exception as e: logger.error(f"Error searching AMM events: {e}", exc_info=True) raise HTTPException(status_code=500, detail=f"Error searching AMM events: {str(e)}") @@ -632,7 +565,7 @@ async def search_amm_positions( status: Optional[str] = None, limit: int = 50, offset: int = 0, - db_manager: AsyncDatabaseManager = Depends(get_database_manager), + amm_service: GatewayAMMService = Depends(get_gateway_amm_service), ): """ Search tracked AMM positions (Meteora DAMM v2 NFTs), newest first. @@ -641,18 +574,10 @@ async def search_amm_positions( come from /gateway/amm/position-info and their history from /gateway/amm/events/search. """ try: - async with db_manager.get_session_context() as session: - repo = GatewayAMMRepository(session) - positions = await repo.search_positions( - connector=connector, network=network, wallet_address=wallet_address, - pool_address=pool_address, status=status, limit=min(limit, 1000), offset=offset, - ) - return { - "data": [repo.position_to_dict(position) for position in positions], - "total_count": len(positions), - "limit": limit, - "offset": offset, - } + return await amm_service.search_positions( + connector=connector, network=network, wallet_address=wallet_address, + pool_address=pool_address, status=status, limit=limit, offset=offset, + ) except Exception as e: logger.error(f"Error searching AMM positions: {e}", exc_info=True) raise HTTPException(status_code=500, detail=f"Error searching AMM positions: {str(e)}") diff --git a/routers/gateway_clmm.py b/routers/gateway_clmm.py index fb0b3858..022a7657 100644 --- a/routers/gateway_clmm.py +++ b/routers/gateway_clmm.py @@ -2,16 +2,13 @@ Gateway CLMM Router - Handles DEX CLMM liquidity operations via Hummingbot Gateway. Supports CLMM connectors (Meteora, Raydium, Uniswap V3) for concentrated liquidity positions. """ -import asyncio import logging from decimal import Decimal from typing import List, Optional from fastapi import APIRouter, Depends, HTTPException, Query -from database import AsyncDatabaseManager -from database.repositories import GatewayCLMMRepository -from deps import get_accounts_service, get_database_manager +from deps import GATEWAY_UNAVAILABLE_DETAIL, get_accounts_service, get_gateway_clmm_service, require_gateway_online from models import ( AMMCreatePoolResponse, CLMMAddLiquidityRequest, @@ -33,12 +30,14 @@ ) from routers.gateway_extras import ( ExtraParamsSpec, + get_transaction_hash_from_response, get_transaction_status_from_response, transaction_id_from_error, validate_extra_params, ) from services.accounts_service import AccountsService from services.gateway_client import GatewayError, check_gateway_error, get_native_gas_token +from services.gateway_clmm_service import GatewayCLMMService logger = logging.getLogger(__name__) @@ -60,154 +59,34 @@ } -async def _refresh_position_data(position, accounts_service: AccountsService, clmm_repo: GatewayCLMMRepository): - """ - Refresh position data from Gateway and update database. - - This updates: - - in_range status - - liquidity amounts - - pending fees - - position status (if closed externally) - """ - try: - # Get wallet address for the position - wallet_address = position.wallet_address - - # Get all positions for this pool and find our specific position - try: - # check_gateway_error is critical here: a Gateway HTTP error must raise (and skip - # the refresh) rather than flow onward and mark the position CLOSED below. - positions_list = check_gateway_error(await accounts_service.gateway_client.clmm_positions_owned( - connector=position.connector, - chain_network=position.network, # position.network is already in 'chain-network' format - wallet_address=wallet_address - )) - - # Find our specific position in the list - result = None - if isinstance(positions_list, list): - for pos in positions_list: - if pos.get("address") == position.position_address: - result = pos - break - - # Absent from a single positions-owned read: could be closed externally, - # could be a lagging RPC node. Closing is owned by the poller's - # consecutive-miss gate (and the zero-liquidity check below) so one - # refresh can never close a live position. - if result is None: - logger.info(f"Position {position.position_address} absent from positions-owned; " - "skipping update (poller's miss-gate owns close detection)") - return - - except Exception as e: - # If we can't fetch positions, log error but don't mark as closed - logger.error(f"Error fetching position from Gateway: {e}") - return - - # Extract current state - current_price = Decimal(str(result.get("price", 0))) - lower_price = Decimal(str(result.get("lowerPrice", 0))) if result.get("lowerPrice") else Decimal("0") - upper_price = Decimal(str(result.get("upperPrice", 0))) if result.get("upperPrice") else Decimal("0") - - # Calculate in_range status - in_range = "UNKNOWN" - if current_price > 0 and lower_price > 0 and upper_price > 0: - if lower_price <= current_price <= upper_price: - in_range = "IN_RANGE" - else: - in_range = "OUT_OF_RANGE" - - # Extract token amounts - base_token_amount = Decimal(str(result.get("baseTokenAmount", 0))) - quote_token_amount = Decimal(str(result.get("quoteTokenAmount", 0))) - - # Check if position has been closed (zero liquidity) - if base_token_amount == 0 and quote_token_amount == 0: - logger.info(f"Position {position.position_address} has zero liquidity, marking as CLOSED") - await clmm_repo.close_position(position.position_address) - return - - # Update liquidity amounts, in_range status, and current price - await clmm_repo.update_position_liquidity( - position_address=position.position_address, - base_token_amount=base_token_amount, - quote_token_amount=quote_token_amount, - in_range=in_range, - current_price=current_price - ) - - # Always write pending fees — 0 is a real value (e.g. right after an - # external collect); the old non-zero guard left stale pendings forever. - base_fee_pending = Decimal(str(result.get("baseFeeAmount", 0))) - quote_fee_pending = Decimal(str(result.get("quoteFeeAmount", 0))) - - await clmm_repo.update_position_fees( - position_address=position.position_address, - base_fee_pending=base_fee_pending, - quote_fee_pending=quote_fee_pending - ) - - logger.debug(f"Refreshed position {position.position_address}: price={current_price}, in_range={in_range}, " - f"base={base_token_amount}, quote={quote_token_amount}") - - except Exception as e: - logger.error(f"Error refreshing position {position.position_address}: {e}", exc_info=True) - raise - - async def _record_failed_write( - db_manager: AsyncDatabaseManager, + clmm_service: GatewayCLMMService, error: Exception, *, event_type: str, position_address: Optional[str], ) -> None: - """Record a write that reached the chain and reverted, before the error is re-raised. + """Pull the transaction id out of a Gateway failure and hand it to the service. - The recording code below only runs when Gateway *returns*. A transaction that landed - and reverted does not return: Gateway raises, the client turns it into a GatewayError, - and control skips every `create_event` call to land in an `except` that persists - nothing. So the database said every operation ever attempted had succeeded, while a - close that reverted at slot 440494812 — costing 0.000011772 SOL — left no row at all. - - Only failures carrying a transaction id are recorded. A pre-flight simulation failure - never got one and cost nothing, and inventing an identifier for it would put a row in - the table that no lookup by hash could ever match. - - Recording never masks the original failure: the caller still gets Gateway's error. + A write that landed on-chain and reverted does not return: Gateway raises, and + the transaction id survives only inside the error message. Parsing it is the + router's job (one parser, shared with the success paths); deciding what row it + becomes is the service's. """ - transaction_hash = transaction_id_from_error(error) - if not transaction_hash or not position_address: - return - - try: - async with db_manager.get_session_context() as session: - repo = GatewayCLMMRepository(session) - position = await repo.get_position_by_address(position_address) - if position is None: - logger.warning( - f"CLMM {event_type} {transaction_hash} reverted on-chain for position " - f"{position_address}, which has no database record — no event written." - ) - return - await repo.create_event({ - "position_id": position.id, - "transaction_hash": transaction_hash, - "event_type": event_type, - "status": "FAILED", - "error_message": str(error), - }) - logger.error( - f"CLMM {event_type} {transaction_hash} landed on-chain and FAILED for position " - f"{position_address}; recorded. {error}" - ) - except Exception as db_error: - logger.error(f"Error recording failed CLMM {event_type}: {db_error}", exc_info=True) - - -@router.get("/clmm/pool-info", response_model=CLMMPoolInfoResponse, response_model_by_alias=False) + await clmm_service.record_failed_write( + transaction_hash=transaction_id_from_error(error), + error=error, + event_type=event_type, + position_address=position_address, + ) + + +@router.get( + "/clmm/pool-info", + response_model=CLMMPoolInfoResponse, + response_model_by_alias=False, + dependencies=[Depends(require_gateway_online)], +) async def get_clmm_pool_info( connector: str, network: str, @@ -236,9 +115,6 @@ async def get_clmm_pool_info( """ try: - if not await accounts_service.gateway_client.ping(): - raise HTTPException(status_code=503, detail="Gateway service is not available") - # Get pool info from Gateway's unified CLMM endpoint result = check_gateway_error(await accounts_service.gateway_client.clmm_pool_info( connector=connector, @@ -424,11 +300,11 @@ async def get_clmm_pools( raise HTTPException(status_code=500, detail=f"Error getting CLMM pools: {str(e)}") -@router.post("/clmm/open", response_model=CLMMOpenPositionResponse) +@router.post("/clmm/open", response_model=CLMMOpenPositionResponse, dependencies=[Depends(require_gateway_online)]) async def open_clmm_position( request: CLMMOpenPositionRequest, accounts_service: AccountsService = Depends(get_accounts_service), - db_manager: AsyncDatabaseManager = Depends(get_database_manager) + clmm_service: GatewayCLMMService = Depends(get_gateway_clmm_service) ): """ Open a NEW CLMM position with initial liquidity. @@ -454,9 +330,6 @@ async def open_clmm_position( validate_extra_params(request.extra_params, CLMM_LIQUIDITY_EXTRA_PARAMS_SPEC, request.connector, "unified /trading/clmm/open") - if not await accounts_service.gateway_client.ping(): - raise HTTPException(status_code=503, detail="Gateway service is not available") - # Parse network_id chain, _ = accounts_service.gateway_client.parse_network_id(request.network) @@ -502,9 +375,9 @@ async def open_clmm_position( extra_params=request.extra_params )) - transaction_hash = result.get("signature") + transaction_hash = get_transaction_hash_from_response(result) if not transaction_hash: - raise HTTPException(status_code=500, detail="No transaction signature returned from Gateway") + raise HTTPException(status_code=500, detail="No transaction hash returned from Gateway") # Gateway's OpenPositionResponse carries position details only inside `data`, # which is present only for CONFIRMED transactions (the response schema strips @@ -562,66 +435,31 @@ async def open_clmm_position( if quote_amount_added is None: quote_amount_added = float(request.quote_token_amount) if request.quote_token_amount else 0 - # Calculate percentage: (upper_price - lower_price) / lower_price - percentage = None - if request.lower_price and request.upper_price and request.lower_price > 0: - percentage = float((request.upper_price - request.lower_price) / request.lower_price) - logger.info(f"Position price range percentage: {percentage:.4f} ({percentage*100:.2f}%)") - # Extract gas fee from Gateway response gas_fee = data.get("fee") gas_token = get_native_gas_token(chain) # Store position and event in database - try: - async with db_manager.get_session_context() as session: - clmm_repo = GatewayCLMMRepository(session) - - # Create position record - position_data = { - "position_address": position_address, - "pool_address": request.pool_address, - "network": request.network, - "connector": request.connector, - "wallet_address": wallet_address, - "trading_pair": trading_pair, - "base_token": base, - "quote_token": quote, - "status": "OPEN", - "lower_price": float(request.lower_price), - "upper_price": float(request.upper_price), - "percentage": percentage, - "entry_price": entry_price, # Pool price when position opened - "current_price": entry_price, # Same as entry at open time, updated by poller - "initial_base_token_amount": float(base_amount_added), - "initial_quote_token_amount": float(quote_amount_added), - "position_rent": float(position_rent) if position_rent else None, - "base_token_amount": float(base_amount_added), - "quote_token_amount": float(quote_amount_added), - "in_range": "UNKNOWN" # Will be updated by poller - } - - position = await clmm_repo.create_position(position_data) - logger.info(f"Recorded CLMM position in database: {position_address}") - - # Create OPEN event with polled status - event_data = { - "position_id": position.id, - "transaction_hash": transaction_hash, - "event_type": "OPEN", - "base_token_amount": float(base_amount_added) if base_amount_added is not None else None, - "quote_token_amount": float(quote_amount_added) if quote_amount_added is not None else None, - "gas_fee": float(gas_fee) if gas_fee is not None else None, - "gas_token": gas_token, - "status": tx_status - } - - await clmm_repo.create_event(event_data) - logger.info(f"Recorded CLMM OPEN event in database: {transaction_hash} " - f"(status: {tx_status}, gas: {gas_fee} {gas_token})") - except Exception as db_error: - # Log but don't fail the operation - it was submitted successfully - logger.error(f"Error recording CLMM position in database: {db_error}", exc_info=True) + await clmm_service.record_open_position( + position_address=position_address, + pool_address=request.pool_address, + network=request.network, + connector=request.connector, + wallet_address=wallet_address, + trading_pair=trading_pair, + base_token=base, + quote_token=quote, + lower_price=request.lower_price, + upper_price=request.upper_price, + entry_price=entry_price, + base_amount_added=base_amount_added, + quote_amount_added=quote_amount_added, + position_rent=position_rent, + transaction_hash=transaction_hash, + gas_fee=gas_fee, + gas_token=gas_token, + tx_status=tx_status, + ) return CLMMOpenPositionResponse( transaction_hash=transaction_hash, @@ -651,11 +489,11 @@ async def open_clmm_position( raise HTTPException(status_code=500, detail=f"Error opening CLMM position: {str(e)}") -@router.post("/clmm/add") +@router.post("/clmm/add", dependencies=[Depends(require_gateway_online)]) async def add_liquidity_to_clmm_position( request: CLMMAddLiquidityRequest, accounts_service: AccountsService = Depends(get_accounts_service), - db_manager: AsyncDatabaseManager = Depends(get_database_manager) + clmm_service: GatewayCLMMService = Depends(get_gateway_clmm_service) ): """ Add MORE liquidity to an EXISTING CLMM position. @@ -677,9 +515,6 @@ async def add_liquidity_to_clmm_position( validate_extra_params(request.extra_params, CLMM_LIQUIDITY_EXTRA_PARAMS_SPEC, request.connector, "unified /trading/clmm/add") - if not await accounts_service.gateway_client.ping(): - raise HTTPException(status_code=503, detail="Gateway service is not available") - # Parse network_id chain, _ = accounts_service.gateway_client.parse_network_id(request.network) @@ -701,7 +536,7 @@ async def add_liquidity_to_clmm_position( extra_params=request.extra_params )) - transaction_hash = result.get("signature") or result.get("txHash") or result.get("hash") + transaction_hash = get_transaction_hash_from_response(result) if not transaction_hash: raise HTTPException(status_code=500, detail="No transaction hash returned from Gateway") @@ -729,15 +564,12 @@ async def add_liquidity_to_clmm_position( # add — the capital is already deposited. add_price = None try: - position_for_pool = None - async with db_manager.get_session_context() as session: - position_for_pool = await GatewayCLMMRepository(session).get_position_by_address( - request.position_address) - if position_for_pool: + pool_address = await clmm_service.get_position_pool_address(request.position_address) + if pool_address: pool_info = check_gateway_error(await accounts_service.gateway_client.clmm_pool_info( connector=request.connector, chain_network=request.network, - pool_address=position_for_pool.pool_address + pool_address=pool_address )) add_price = float(pool_info.get("price")) if pool_info.get("price") else None except Exception as price_error: @@ -745,46 +577,16 @@ async def add_liquidity_to_clmm_position( f"entry price will not be re-weighted: {price_error}") # Store ADD_LIQUIDITY event in database - try: - async with db_manager.get_session_context() as session: - clmm_repo = GatewayCLMMRepository(session) - - # Get position to link event - position = await clmm_repo.get_position_by_address(request.position_address) - if position: - event_data = { - "position_id": position.id, - "transaction_hash": transaction_hash, - "event_type": "ADD_LIQUIDITY", - # `is not None`: 0 is a real amount on single-sided adds - "base_token_amount": float(base_amount_added) if base_amount_added is not None else None, - "quote_token_amount": float(quote_amount_added) if quote_amount_added is not None else None, - "gas_fee": float(gas_fee) if gas_fee is not None else None, - "gas_token": gas_token, - "status": tx_status - } - await clmm_repo.create_event(event_data) - logger.info(f"Recorded CLMM ADD_LIQUIDITY event: {transaction_hash} " - f"(status: {tx_status}, gas: {gas_fee} {gas_token})") - - # Added capital raises both the PnL baseline and the held amounts. - # Book here only when the tx confirmed inline (the event is created - # CONFIRMED and the poller never re-processes it); SUBMITTED events - # are booked by the poller's confirm path. - if tx_status == "CONFIRMED": - await clmm_repo.add_to_position_amounts( - position_address=request.position_address, - base_delta=Decimal(str(base_amount_added or 0)), - quote_delta=Decimal(str(quote_amount_added or 0)), - entry_price=Decimal(str(add_price)) if add_price else None, - ) - else: - logger.warning(f"ADD_LIQUIDITY {transaction_hash} executed for position " - f"{request.position_address} with no database record — " - "no event recorded (position may be a pending open " - "not yet discovered)") - except Exception as db_error: - logger.error(f"Error recording ADD_LIQUIDITY event: {db_error}", exc_info=True) + await clmm_service.record_add_liquidity( + position_address=request.position_address, + transaction_hash=transaction_hash, + tx_status=tx_status, + base_amount_added=base_amount_added, + quote_amount_added=quote_amount_added, + gas_fee=gas_fee, + gas_token=gas_token, + add_price=add_price, + ) return { "transaction_hash": transaction_hash, @@ -799,7 +601,7 @@ async def add_liquidity_to_clmm_position( raise except GatewayError as e: await _record_failed_write( - db_manager, e, event_type="ADD_LIQUIDITY", position_address=request.position_address + clmm_service, e, event_type="ADD_LIQUIDITY", position_address=request.position_address ) raise HTTPException(status_code=e.status, detail=f"Gateway error adding liquidity: {e}") except ValueError as e: @@ -809,11 +611,11 @@ async def add_liquidity_to_clmm_position( raise HTTPException(status_code=500, detail=f"Error adding liquidity to CLMM position: {str(e)}") -@router.post("/clmm/remove") +@router.post("/clmm/remove", dependencies=[Depends(require_gateway_online)]) async def remove_liquidity_from_clmm_position( request: CLMMRemoveLiquidityRequest, accounts_service: AccountsService = Depends(get_accounts_service), - db_manager: AsyncDatabaseManager = Depends(get_database_manager) + clmm_service: GatewayCLMMService = Depends(get_gateway_clmm_service) ): """ Remove SOME liquidity from a CLMM position (partial removal). @@ -830,9 +632,6 @@ async def remove_liquidity_from_clmm_position( Transaction hash """ try: - if not await accounts_service.gateway_client.ping(): - raise HTTPException(status_code=503, detail="Gateway service is not available") - # Parse network_id chain, _ = accounts_service.gateway_client.parse_network_id(request.network) @@ -852,7 +651,7 @@ async def remove_liquidity_from_clmm_position( slippage_pct=float(request.slippage_pct) if request.slippage_pct is not None else None )) - transaction_hash = result.get("signature") or result.get("txHash") or result.get("hash") + transaction_hash = get_transaction_hash_from_response(result) if not transaction_hash: raise HTTPException(status_code=500, detail="No transaction hash returned from Gateway") @@ -870,47 +669,15 @@ async def remove_liquidity_from_clmm_position( quote_amount_removed = data.get("quoteTokenAmountRemoved") # Store REMOVE_LIQUIDITY event in database - try: - async with db_manager.get_session_context() as session: - clmm_repo = GatewayCLMMRepository(session) - - # Get position to link event - position = await clmm_repo.get_position_by_address(request.position_address) - if position: - # No "percentage" key: GatewayCLMMEvent has no such column, and the - # stray kwarg made create_event raise — silently losing every - # REMOVE_LIQUIDITY event to the log-and-continue handler below. - event_data = { - "position_id": position.id, - "transaction_hash": transaction_hash, - "event_type": "REMOVE_LIQUIDITY", - "base_token_amount": float(base_amount_removed) if base_amount_removed is not None else None, - "quote_token_amount": float(quote_amount_removed) if quote_amount_removed is not None else None, - "gas_fee": float(gas_fee) if gas_fee is not None else None, - "gas_token": gas_token, - "status": tx_status - } - await clmm_repo.create_event(event_data) - logger.info(f"Recorded CLMM REMOVE_LIQUIDITY event: {transaction_hash} " - f"(status: {tx_status}, gas: {gas_fee} {gas_token})") - - # Withdrawn capital lowers both the held amounts and the PnL - # baseline. Book only on inline confirmation (the event is created - # CONFIRMED and the poller never re-processes it); SUBMITTED events - # are booked by the poller's confirm path. - if tx_status == "CONFIRMED": - await clmm_repo.subtract_from_position_amounts( - position_address=request.position_address, - base_delta=Decimal(str(base_amount_removed or 0)), - quote_delta=Decimal(str(quote_amount_removed or 0)), - ) - else: - logger.warning(f"REMOVE_LIQUIDITY {transaction_hash} executed for position " - f"{request.position_address} with no database record — " - "no event recorded (position may be a pending open " - "not yet discovered)") - except Exception as db_error: - logger.error(f"Error recording REMOVE_LIQUIDITY event: {db_error}", exc_info=True) + await clmm_service.record_remove_liquidity( + position_address=request.position_address, + transaction_hash=transaction_hash, + tx_status=tx_status, + base_amount_removed=base_amount_removed, + quote_amount_removed=quote_amount_removed, + gas_fee=gas_fee, + gas_token=gas_token, + ) return { "transaction_hash": transaction_hash, @@ -926,7 +693,7 @@ async def remove_liquidity_from_clmm_position( raise except GatewayError as e: await _record_failed_write( - db_manager, e, event_type="REMOVE_LIQUIDITY", position_address=request.position_address + clmm_service, e, event_type="REMOVE_LIQUIDITY", position_address=request.position_address ) raise HTTPException(status_code=e.status, detail=f"Gateway error removing liquidity: {e}") except ValueError as e: @@ -936,11 +703,11 @@ async def remove_liquidity_from_clmm_position( raise HTTPException(status_code=500, detail=f"Error removing liquidity from CLMM position: {str(e)}") -@router.post("/clmm/close", response_model=CLMMClosePositionResponse) +@router.post("/clmm/close", response_model=CLMMClosePositionResponse, dependencies=[Depends(require_gateway_online)]) async def close_clmm_position( request: CLMMClosePositionRequest, accounts_service: AccountsService = Depends(get_accounts_service), - db_manager: AsyncDatabaseManager = Depends(get_database_manager) + clmm_service: GatewayCLMMService = Depends(get_gateway_clmm_service) ): """ CLOSE a CLMM position completely (removes all liquidity and collects pending fees). @@ -955,20 +722,12 @@ async def close_clmm_position( Transaction hash and collected fee amounts """ try: - if not await accounts_service.gateway_client.ping(): - raise HTTPException(status_code=503, detail="Gateway service is not available") - # Parse network_id chain, _ = accounts_service.gateway_client.parse_network_id(request.network) # Wallet resolution: an explicit request value wins (same precedence as # open/add/remove), then the DB row's wallet, then the default wallet. - db_wallet = None - async with db_manager.get_session_context() as session: - clmm_repo = GatewayCLMMRepository(session) - db_position = await clmm_repo.get_position_by_address(request.position_address) - if db_position: - db_wallet = db_position.wallet_address + db_wallet = await clmm_service.get_position_wallet(request.position_address) wallet_address = request.wallet_address or db_wallet wallet_address = await accounts_service.gateway_client.get_wallet_address_or_default( @@ -1015,7 +774,7 @@ async def close_clmm_position( slippage_pct=float(request.slippage_pct) if request.slippage_pct is not None else None, )) - transaction_hash = result.get("signature") or result.get("txHash") or result.get("hash") + transaction_hash = get_transaction_hash_from_response(result) if not transaction_hash: raise HTTPException(status_code=500, detail="No transaction hash returned from Gateway") @@ -1048,104 +807,22 @@ async def close_clmm_position( f"rent refunded={position_rent_refunded}") # Store CLOSE event in database and update position - try: - async with db_manager.get_session_context() as session: - clmm_repo = GatewayCLMMRepository(session) - - # Get position to link event - position = await clmm_repo.get_position_by_address(request.position_address) - if position: - # Create event record - event_data = { - "position_id": position.id, - "transaction_hash": transaction_hash, - "event_type": "CLOSE", - "base_token_amount": float(base_amount_removed) if base_amount_removed is not None else None, - "quote_token_amount": float(quote_amount_removed) if quote_amount_removed is not None else None, - "base_fee_collected": float(base_fee_collected) if base_fee_collected is not None else None, - "quote_fee_collected": float(quote_fee_collected) if quote_fee_collected is not None else None, - "gas_fee": float(gas_fee) if gas_fee is not None else None, - "gas_token": gas_token, - "status": tx_status - } - await clmm_repo.create_event(event_data) - logger.info(f"Recorded CLMM CLOSE event: {transaction_hash} " - f"(status: {tx_status}, gas: {gas_fee} {gas_token})") - - # Position bookkeeping happens exactly once, when the tx is known - # good: CONFIRMED here (the event is created CONFIRMED, so the - # poller never touches it), or in the poller's confirm path for - # SUBMITTED events. A FAILED tx mutates nothing — the old - # unconditional booking permanently inflated *_fee_collected on - # failed closes. - if tx_status == "CONFIRMED": - new_base_collected = Decimal(str(position.base_fee_collected)) + base_fee_collected - new_quote_collected = Decimal(str(position.quote_fee_collected)) + quote_fee_collected - - await clmm_repo.update_position_fees( - position_address=request.position_address, - base_fee_collected=new_base_collected, - quote_fee_collected=new_quote_collected, - base_fee_pending=Decimal("0"), - quote_fee_pending=Decimal("0") - ) - - # Update current_price with close price - if close_price: - await clmm_repo.update_position_liquidity( - position_address=request.position_address, - base_token_amount=Decimal(str(position.base_token_amount)), - quote_token_amount=Decimal(str(position.quote_token_amount)), - current_price=Decimal(str(close_price)) - ) - - # Verify position is actually gone on Gateway before marking - # CLOSED (some connectors 500 instead of 404 for a - # nonexistent position — right after our own close, either - # means gone). - try: - await asyncio.sleep(2) # Wait for transaction to propagate - - verify_result = await accounts_service.gateway_client.clmm_position_info( - connector=request.connector, - chain_network=request.network, - position_address=request.position_address - ) - - if verify_result and isinstance(verify_result, dict) and "error" in verify_result: - status_code = verify_result.get("status") - if status_code in (404, 500): - await clmm_repo.close_position( - request.position_address, - position_rent_refunded=(Decimal(str(position_rent_refunded)) - if position_rent_refunded is not None else None) - ) - logger.info(f"Position {request.position_address} verified as closed " - f"(Gateway returned {status_code})") - else: - logger.warning(f"Unexpected error verifying position close: {verify_result}") - elif verify_result and "address" in verify_result: - # Position still exists - might be a failed close or delayed propagation - logger.warning(f"Position {request.position_address} still exists after close " - "transaction. Will be handled by poller.") - else: - logger.debug("Could not verify position close status, will be handled by poller") - - except Exception as verify_error: - logger.warning(f"Error verifying position close: {verify_error}. Will be handled by poller.") - - logger.info(f"Updated position {request.position_address}: " - "collected fees updated, pending fees reset to 0.") - else: - # H8 window: a close on a position hapi has no row for (e.g. a - # pending open awaiting the discovery sweep) leaves no event — - # say so loudly instead of silently skipping. - logger.warning(f"CLOSE {transaction_hash} executed for position " - f"{request.position_address} with no database record — " - "no CLOSE event recorded (position may be a pending open " - "not yet discovered)") - except Exception as db_error: - logger.error(f"Error recording CLOSE event: {db_error}", exc_info=True) + await clmm_service.record_close( + position_address=request.position_address, + connector=request.connector, + network=request.network, + transaction_hash=transaction_hash, + tx_status=tx_status, + base_amount_removed=base_amount_removed, + quote_amount_removed=quote_amount_removed, + base_fee_collected=base_fee_collected, + quote_fee_collected=quote_fee_collected, + position_rent_refunded=position_rent_refunded, + close_price=close_price, + gas_fee=gas_fee, + gas_token=gas_token, + gateway_client=accounts_service.gateway_client, + ) return CLMMClosePositionResponse( transaction_hash=transaction_hash, @@ -1162,7 +839,7 @@ async def close_clmm_position( raise except GatewayError as e: await _record_failed_write( - db_manager, e, event_type="CLOSE", position_address=request.position_address + clmm_service, e, event_type="CLOSE", position_address=request.position_address ) raise HTTPException(status_code=e.status, detail=f"Gateway error closing CLMM position: {e}") except ValueError as e: @@ -1172,11 +849,11 @@ async def close_clmm_position( raise HTTPException(status_code=500, detail=f"Error closing CLMM position: {str(e)}") -@router.post("/clmm/collect-fees", response_model=CLMMCollectFeesResponse) +@router.post("/clmm/collect-fees", response_model=CLMMCollectFeesResponse, dependencies=[Depends(require_gateway_online)]) async def collect_fees_from_clmm_position( request: CLMMCollectFeesRequest, accounts_service: AccountsService = Depends(get_accounts_service), - db_manager: AsyncDatabaseManager = Depends(get_database_manager) + clmm_service: GatewayCLMMService = Depends(get_gateway_clmm_service) ): """ Collect accumulated fees from a CLMM liquidity position. @@ -1191,20 +868,12 @@ async def collect_fees_from_clmm_position( Transaction hash and collected fee amounts """ try: - if not await accounts_service.gateway_client.ping(): - raise HTTPException(status_code=503, detail="Gateway service is not available") - # Parse network_id chain, _ = accounts_service.gateway_client.parse_network_id(request.network) # Wallet resolution: an explicit request value wins (same precedence as # open/add/remove), then the DB row's wallet, then the default wallet. - db_wallet = None - async with db_manager.get_session_context() as session: - clmm_repo = GatewayCLMMRepository(session) - db_position = await clmm_repo.get_position_by_address(request.position_address) - if db_position: - db_wallet = db_position.wallet_address + db_wallet = await clmm_service.get_position_wallet(request.position_address) wallet_address = request.wallet_address or db_wallet wallet_address = await accounts_service.gateway_client.get_wallet_address_or_default( @@ -1247,7 +916,7 @@ async def collect_fees_from_clmm_position( position_address=request.position_address )) - transaction_hash = result.get("signature") or result.get("txHash") or result.get("hash") + transaction_hash = get_transaction_hash_from_response(result) if not transaction_hash: raise HTTPException(status_code=500, detail="No transaction hash returned from Gateway") @@ -1271,52 +940,15 @@ async def collect_fees_from_clmm_position( logger.info(f"Collected fees: base={base_fee_collected}, quote={quote_fee_collected}") # Store COLLECT_FEES event in database and update position - try: - async with db_manager.get_session_context() as session: - clmm_repo = GatewayCLMMRepository(session) - - # Get position to link event - position = await clmm_repo.get_position_by_address(request.position_address) - if position: - # Create event record - event_data = { - "position_id": position.id, - "transaction_hash": transaction_hash, - "event_type": "COLLECT_FEES", - "base_fee_collected": float(base_fee_collected) if base_fee_collected is not None else None, - "quote_fee_collected": float(quote_fee_collected) if quote_fee_collected is not None else None, - "gas_fee": float(gas_fee) if gas_fee is not None else None, - "gas_token": gas_token, - "status": tx_status - } - await clmm_repo.create_event(event_data) - logger.info(f"Recorded CLMM COLLECT_FEES event: {transaction_hash} " - f"(status: {tx_status}, gas: {gas_fee} {gas_token})") - - # Book fees exactly once: CONFIRMED here (event created CONFIRMED, - # never re-processed), SUBMITTED in the poller's confirm path. - # The old unconditional booking double-counted every pending - # collect (endpoint + poller) and kept phantom fees on failures. - if tx_status == "CONFIRMED": - new_base_collected = Decimal(str(position.base_fee_collected)) + base_fee_collected - new_quote_collected = Decimal(str(position.quote_fee_collected)) + quote_fee_collected - - await clmm_repo.update_position_fees( - position_address=request.position_address, - base_fee_collected=new_base_collected, - quote_fee_collected=new_quote_collected, - base_fee_pending=Decimal("0"), - quote_fee_pending=Decimal("0") - ) - logger.info(f"Updated position {request.position_address}: " - "collected fees updated, pending fees reset to 0") - else: - logger.warning(f"COLLECT_FEES {transaction_hash} executed for position " - f"{request.position_address} with no database record — " - "no event recorded (position may be a pending open " - "not yet discovered)") - except Exception as db_error: - logger.error(f"Error recording COLLECT_FEES event: {db_error}", exc_info=True) + await clmm_service.record_collect_fees( + position_address=request.position_address, + transaction_hash=transaction_hash, + tx_status=tx_status, + base_fee_collected=base_fee_collected, + quote_fee_collected=quote_fee_collected, + gas_fee=gas_fee, + gas_token=gas_token, + ) return CLMMCollectFeesResponse( transaction_hash=transaction_hash, @@ -1330,7 +962,7 @@ async def collect_fees_from_clmm_position( raise except GatewayError as e: await _record_failed_write( - db_manager, e, event_type="COLLECT_FEES", position_address=request.position_address + clmm_service, e, event_type="COLLECT_FEES", position_address=request.position_address ) raise HTTPException(status_code=e.status, detail=f"Gateway error collecting fees: {e}") except ValueError as e: @@ -1340,7 +972,7 @@ async def collect_fees_from_clmm_position( raise HTTPException(status_code=500, detail=f"Error collecting fees: {str(e)}") -@router.post("/clmm/positions_owned", response_model=List[CLMMPositionInfo]) +@router.post("/clmm/positions_owned", response_model=List[CLMMPositionInfo], dependencies=[Depends(require_gateway_online)]) async def get_clmm_positions_owned( request: CLMMPositionsOwnedRequest, accounts_service: AccountsService = Depends(get_accounts_service) @@ -1362,9 +994,6 @@ async def get_clmm_positions_owned( List of CLMM position information """ try: - if not await accounts_service.gateway_client.ping(): - raise HTTPException(status_code=503, detail="Gateway service is not available") - # Parse network_id chain, network = accounts_service.gateway_client.parse_network_id(request.network) @@ -1437,7 +1066,12 @@ async def get_clmm_positions_owned( raise HTTPException(status_code=500, detail=f"Error getting CLMM positions owned: {str(e)}") -@router.post("/clmm/quote-position", response_model=CLMMQuotePositionResponse, response_model_by_alias=False) +@router.post( + "/clmm/quote-position", + response_model=CLMMQuotePositionResponse, + response_model_by_alias=False, + dependencies=[Depends(require_gateway_online)], +) async def quote_clmm_position( request: CLMMQuotePositionRequest, accounts_service: AccountsService = Depends(get_accounts_service) @@ -1450,9 +1084,6 @@ async def quote_clmm_position( (and which side limits it), without signing or submitting anything. """ try: - if not await accounts_service.gateway_client.ping(): - raise HTTPException(status_code=503, detail="Gateway service is not available") - result = check_gateway_error(await accounts_service.gateway_client.clmm_quote_position( connector=request.connector, chain_network=request.network, @@ -1476,7 +1107,12 @@ async def quote_clmm_position( raise HTTPException(status_code=500, detail=f"Error quoting CLMM position: {str(e)}") -@router.post("/clmm/create-pool", response_model=AMMCreatePoolResponse, response_model_by_alias=False) +@router.post( + "/clmm/create-pool", + response_model=AMMCreatePoolResponse, + response_model_by_alias=False, + dependencies=[Depends(require_gateway_online)], +) async def create_clmm_pool( request: CLMMCreatePoolRequest, accounts_service: AccountsService = Depends(get_accounts_service) @@ -1492,9 +1128,6 @@ async def create_clmm_pool( validate_extra_params(request.extra_params, CLMM_CREATE_POOL_EXTRA_PARAMS_SPEC, request.connector, "unified /trading/clmm/create-pool") - if not await accounts_service.gateway_client.ping(): - raise HTTPException(status_code=503, detail="Gateway service is not available") - chain, _ = accounts_service.gateway_client.parse_network_id(request.network) wallet_address = await accounts_service.gateway_client.get_wallet_address_or_default( chain=chain, @@ -1526,7 +1159,7 @@ async def create_clmm_pool( raise HTTPException(status_code=500, detail=f"Error creating CLMM pool: {str(e)}") -@router.get("/clmm/position-info", response_model=CLMMPositionInfo) +@router.get("/clmm/position-info", response_model=CLMMPositionInfo, dependencies=[Depends(require_gateway_online)]) async def get_clmm_position_info( connector: str, network: str, @@ -1540,9 +1173,6 @@ async def get_clmm_position_info( or closed position as an error (500/404), surfaced here as 404. """ try: - if not await accounts_service.gateway_client.ping(): - raise HTTPException(status_code=503, detail="Gateway service is not available") - pos = await accounts_service.gateway_client.clmm_position_info( connector=connector, chain_network=network, @@ -1550,7 +1180,7 @@ async def get_clmm_position_info( ) if pos is None or not isinstance(pos, dict): # Connection error: the client returns None — a 503, not a crash. - raise HTTPException(status_code=503, detail="Gateway service is not available") + raise HTTPException(status_code=503, detail=GATEWAY_UNAVAILABLE_DETAIL) if "error" in pos: status_code = pos.get("status") if status_code in (404, 500): @@ -1600,7 +1230,7 @@ async def get_clmm_position_events( position_address: str, event_type: Optional[str] = None, limit: int = 100, - db_manager: AsyncDatabaseManager = Depends(get_database_manager) + clmm_service: GatewayCLMMService = Depends(get_gateway_clmm_service) ): """ Get event history for a CLMM position. @@ -1616,18 +1246,11 @@ async def get_clmm_position_events( List of position events """ try: - async with db_manager.get_session_context() as session: - clmm_repo = GatewayCLMMRepository(session) - events = await clmm_repo.get_position_events( - position_address=position_address, - event_type=event_type, - limit=limit - ) - - return { - "data": [clmm_repo.event_to_dict(event) for event in events], - "total_count": len(events) - } + return await clmm_service.get_position_events( + position_address=position_address, + event_type=event_type, + limit=limit + ) except Exception as e: logger.error(f"Error getting position events: {e}", exc_info=True) @@ -1645,7 +1268,7 @@ async def search_clmm_positions( limit: int = 50, offset: int = 0, refresh: bool = False, - db_manager: AsyncDatabaseManager = Depends(get_database_manager), + clmm_service: GatewayCLMMService = Depends(get_gateway_clmm_service), accounts_service: AccountsService = Depends(get_accounts_service) ): """ @@ -1667,78 +1290,18 @@ async def search_clmm_positions( Paginated list of positions """ try: - # Validate limit - if limit > 1000: - limit = 1000 - - # Optionally refresh position data from Gateway first - if refresh and await accounts_service.gateway_client.ping(): - # Get positions to refresh - async with db_manager.get_session_context() as session: - clmm_repo = GatewayCLMMRepository(session) - positions_to_refresh = await clmm_repo.get_positions( - network=network, - connector=connector, - wallet_address=wallet_address, - trading_pair=trading_pair, - status=status, - position_addresses=position_addresses, - limit=limit, - offset=offset - ) - - # Extract position addresses and details before closing session - position_details = [ - { - "position_address": pos.position_address, - "pool_address": pos.pool_address, - "connector": pos.connector, - "network": pos.network, - "wallet_address": pos.wallet_address - } - for pos in positions_to_refresh - ] - - # Refresh each position in a separate session - logger.info(f"Refreshing {len(position_details)} positions from Gateway") - for pos_detail in position_details: - try: - async with db_manager.get_session_context() as session: - clmm_repo = GatewayCLMMRepository(session) - # Get position again in this session - position = await clmm_repo.get_position_by_address(pos_detail["position_address"]) - if position: - await _refresh_position_data(position, accounts_service, clmm_repo) - except Exception as e: - logger.warning(f"Failed to refresh position {pos_detail['position_address']}: {e}") - # Continue with other positions even if one fails - - # Get final results after refresh - async with db_manager.get_session_context() as session: - clmm_repo = GatewayCLMMRepository(session) - positions = await clmm_repo.get_positions( - network=network, - connector=connector, - wallet_address=wallet_address, - trading_pair=trading_pair, - status=status, - position_addresses=position_addresses, - limit=limit, - offset=offset - ) - - # Get total count for pagination - has_more = len(positions) == limit - - return { - "data": [clmm_repo.position_to_dict(pos) for pos in positions], - "pagination": { - "limit": limit, - "offset": offset, - "has_more": has_more, - "total_count": len(positions) + offset if not has_more else None - } - } + return await clmm_service.search_positions( + network=network, + connector=connector, + wallet_address=wallet_address, + trading_pair=trading_pair, + status=status, + position_addresses=position_addresses, + limit=limit, + offset=offset, + refresh=refresh, + gateway_client=accounts_service.gateway_client, + ) except Exception as e: logger.error(f"Error searching CLMM positions: {e}", exc_info=True) diff --git a/routers/gateway_extras.py b/routers/gateway_extras.py index 79ec660d..a6d7262a 100644 --- a/routers/gateway_extras.py +++ b/routers/gateway_extras.py @@ -33,6 +33,26 @@ def get_transaction_status_from_response(gateway_response: dict) -> str: return "SUBMITTED" +def get_transaction_hash_from_response(gateway_response: dict) -> Optional[str]: + """The transaction id a Gateway write response reports, whatever chain it came from. + + Gateway names the same field three ways — Solana routes answer `signature`, EVM + routes `txHash`, and a few older shapes `hash` — so reading only one of them is + reading only one family of chains. The CLMM open handler did exactly that + (`signature` alone) and answered every uniswap/pancakeswap open with a 500 saying + no signature was returned, for a position that had in fact just been opened. + + Returns None when the response carries no id at all; callers decide whether that is + a 500 (the routes that must answer with a hash) or an empty column (the AMM event + recorder, which records what it has and never fails the write over it). + """ + return ( + gateway_response.get("signature") + or gateway_response.get("txHash") + or gateway_response.get("hash") + ) + + # A Solana signature or an EVM transaction hash, as they appear inside Gateway's # landed-but-failed message: "Transaction landed on-chain but failed: ". _TRANSACTION_ID = re.compile(r"[Tt]ransaction ([1-9A-HJ-NP-Za-km-z]{43,88}|0x[0-9a-fA-F]{64})") diff --git a/routers/gateway_swap.py b/routers/gateway_swap.py index b2f7bd85..a5eba64e 100644 --- a/routers/gateway_swap.py +++ b/routers/gateway_swap.py @@ -11,13 +11,17 @@ from fastapi import APIRouter, Depends, HTTPException -from database import AsyncDatabaseManager -from database.repositories import GatewaySwapRepository -from deps import get_accounts_service, get_database_manager +from deps import get_accounts_service, get_gateway_swap_service, require_gateway_online from models import SwapExecuteQuoteRequest, SwapExecuteRequest, SwapExecuteResponse, SwapQuoteRequest, SwapQuoteResponse -from routers.gateway_extras import ExtraParamsSpec, get_transaction_status_from_response, validate_extra_params +from routers.gateway_extras import ( + ExtraParamsSpec, + get_transaction_hash_from_response, + get_transaction_status_from_response, + validate_extra_params, +) from services.accounts_service import AccountsService from services.gateway_client import GatewayError, check_gateway_error, get_native_gas_token +from services.gateway_swap_service import GatewaySwapService from utils.trading_pair import split_trading_pair logger = logging.getLogger(__name__) @@ -32,7 +36,7 @@ } -@router.post("/swap/quote", response_model=SwapQuoteResponse) +@router.post("/swap/quote", response_model=SwapQuoteResponse, dependencies=[Depends(require_gateway_online)]) async def get_swap_quote( request: SwapQuoteRequest, accounts_service: AccountsService = Depends(get_accounts_service) @@ -56,9 +60,6 @@ async def get_swap_quote( validate_extra_params(request.extra_params, SWAP_EXTRA_PARAMS_SPEC, request.connector, "/trading/{router,clmm,amm}/quote-swap") - if not await accounts_service.gateway_client.ping(): - raise HTTPException(status_code=503, detail="Gateway service is not available") - # Parse trading pair base, quote = split_trading_pair(request.trading_pair) @@ -119,7 +120,7 @@ def _dec(key): async def _record_and_report_swap( *, result: dict, - db_manager: AsyncDatabaseManager, + swap_service: GatewaySwapService, accounts_service: AccountsService, connector: str, network: str, @@ -138,7 +139,7 @@ async def _record_and_report_swap( already saw — and not at all in what has to be recorded afterwards. One copy is what stops the two-step flow from growing its own subtly different accounting. """ - transaction_hash = result.get("signature") or result.get("txHash") or result.get("hash") + transaction_hash = get_transaction_hash_from_response(result) if not transaction_hash: raise HTTPException(status_code=500, detail="No transaction hash returned from Gateway") @@ -192,39 +193,26 @@ async def _record_and_report_swap( # Get transaction status from Gateway response tx_status = get_transaction_status_from_response(result) - # Store swap in database - try: - async with db_manager.get_session_context() as session: - swap_repo = GatewaySwapRepository(session) - - swap_data = { - "transaction_hash": transaction_hash, - "network": network, - # Store the base venue name: a swap on "jupiter" and one on - # "jupiter/router" are the same venue and must file together. - "connector": connector.split("/")[0], - "wallet_address": wallet_address, - "trading_pair": trading_pair, - "base_token": base, - "quote_token": quote, - "side": side, - "input_amount": float(input_amount), - "output_amount": float(output_amount), - "price": float(price), - "slippage_pct": float(slippage_pct) if slippage_pct is not None else None, - "gas_fee": float(gas_fee) if gas_fee is not None else None, - "gas_token": gas_token, - "status": tx_status, - # Set by the pool-scoped routes, which resolve exactly one pool; a - # router picks its own path across pools and leaves it unset. - "pool_address": data.get("poolAddress") - } - - await swap_repo.create_swap(swap_data) - logger.info(f"Recorded swap in database: {transaction_hash} (status: {tx_status})") - except Exception as db_error: - # Log but don't fail the swap - it was submitted successfully - logger.error(f"Error recording swap in database: {db_error}", exc_info=True) + # Store swap in database. Best-effort by policy — the swap was submitted + # successfully, so a bookkeeping failure must not be reported as a failed swap. + await swap_service.record_swap( + transaction_hash=transaction_hash, + network=network, + connector=connector, + wallet_address=wallet_address, + trading_pair=trading_pair, + base_token=base, + quote_token=quote, + side=side, + input_amount=input_amount, + output_amount=output_amount, + price=price, + slippage_pct=slippage_pct, + gas_fee=gas_fee, + gas_token=gas_token, + status=tx_status, + pool_address=data.get("poolAddress"), + ) return SwapExecuteResponse( transaction_hash=transaction_hash, @@ -242,11 +230,11 @@ async def _record_and_report_swap( ) -@router.post("/swap/execute", response_model=SwapExecuteResponse) +@router.post("/swap/execute", response_model=SwapExecuteResponse, dependencies=[Depends(require_gateway_online)]) async def execute_swap( request: SwapExecuteRequest, accounts_service: AccountsService = Depends(get_accounts_service), - db_manager: AsyncDatabaseManager = Depends(get_database_manager) + swap_service: GatewaySwapService = Depends(get_gateway_swap_service) ): """ Execute a swap transaction via router (Jupiter, 0x). @@ -268,9 +256,6 @@ async def execute_swap( validate_extra_params(request.extra_params, SWAP_EXTRA_PARAMS_SPEC, request.connector, "/trading/{router,clmm,amm}/execute-swap") - if not await accounts_service.gateway_client.ping(): - raise HTTPException(status_code=503, detail="Gateway service is not available") - # Parse network_id (chain needed for wallet lookup) chain, network = accounts_service.gateway_client.parse_network_id(request.network) @@ -297,7 +282,7 @@ async def execute_swap( )) return await _record_and_report_swap( result=result, - db_manager=db_manager, + swap_service=swap_service, accounts_service=accounts_service, connector=request.connector, network=request.network, @@ -324,7 +309,7 @@ async def execute_swap( @router.get("/swaps/{transaction_hash}/status") async def get_swap_status( transaction_hash: str, - db_manager: AsyncDatabaseManager = Depends(get_database_manager) + swap_service: GatewaySwapService = Depends(get_gateway_swap_service) ): """ Get status of a specific swap by transaction hash. @@ -336,14 +321,12 @@ async def get_swap_status( Swap details including current status """ try: - async with db_manager.get_session_context() as session: - swap_repo = GatewaySwapRepository(session) - swap = await swap_repo.get_swap_by_tx_hash(transaction_hash) - - if not swap: - raise HTTPException(status_code=404, detail=f"Swap not found: {transaction_hash}") - - return swap_repo.to_dict(swap) + # None means "no such row" and nothing else: the service lets a database + # failure raise, so an unreachable database can never read as a 404. + swap = await swap_service.get_swap(transaction_hash) + if swap is None: + raise HTTPException(status_code=404, detail=f"Swap not found: {transaction_hash}") + return swap except HTTPException: raise @@ -352,11 +335,11 @@ async def get_swap_status( raise HTTPException(status_code=500, detail=f"Error getting swap status: {str(e)}") -@router.post("/swap/execute-quote", response_model=SwapExecuteResponse) +@router.post("/swap/execute-quote", response_model=SwapExecuteResponse, dependencies=[Depends(require_gateway_online)]) async def execute_swap_quote( request: SwapExecuteQuoteRequest, accounts_service: AccountsService = Depends(get_accounts_service), - db_manager: AsyncDatabaseManager = Depends(get_database_manager) + swap_service: GatewaySwapService = Depends(get_gateway_swap_service) ): """ Execute a quote returned by /swap/quote, by its quote_id. @@ -378,9 +361,6 @@ async def execute_swap_quote( Transaction hash and what the swap actually moved """ try: - if not await accounts_service.gateway_client.ping(): - raise HTTPException(status_code=503, detail="Gateway service is not available") - chain, network = accounts_service.gateway_client.parse_network_id(request.network) wallet_address = await accounts_service.gateway_client.get_wallet_address_or_default( chain=chain, @@ -402,7 +382,7 @@ async def execute_swap_quote( return await _record_and_report_swap( result=result, - db_manager=db_manager, + swap_service=swap_service, accounts_service=accounts_service, connector=request.connector, network=request.network, @@ -437,7 +417,7 @@ async def search_swaps( end_time: Optional[int] = None, limit: int = 50, offset: int = 0, - db_manager: AsyncDatabaseManager = Depends(get_database_manager) + swap_service: GatewaySwapService = Depends(get_gateway_swap_service) ): """ Search swap history with filters. @@ -457,36 +437,17 @@ async def search_swaps( Paginated list of swaps """ try: - # Validate limit - if limit > 1000: - limit = 1000 - - async with db_manager.get_session_context() as session: - swap_repo = GatewaySwapRepository(session) - swaps = await swap_repo.get_swaps( - network=network, - connector=connector, - wallet_address=wallet_address, - trading_pair=trading_pair, - status=status, - start_time=start_time, - end_time=end_time, - limit=limit, - offset=offset - ) - - # Get total count for pagination (simplified - actual count would need separate query) - has_more = len(swaps) == limit - - return { - "data": [swap_repo.to_dict(swap) for swap in swaps], - "pagination": { - "limit": limit, - "offset": offset, - "has_more": has_more, - "total_count": len(swaps) + offset if not has_more else None - } - } + return await swap_service.search_swaps( + network=network, + connector=connector, + wallet_address=wallet_address, + trading_pair=trading_pair, + status=status, + start_time=start_time, + end_time=end_time, + limit=limit, + offset=offset + ) except Exception as e: logger.error(f"Error searching swaps: {e}", exc_info=True) @@ -499,7 +460,7 @@ async def get_swaps_summary( wallet_address: Optional[str] = None, start_time: Optional[int] = None, end_time: Optional[int] = None, - db_manager: AsyncDatabaseManager = Depends(get_database_manager) + swap_service: GatewaySwapService = Depends(get_gateway_swap_service) ): """ Get swap summary statistics. @@ -514,15 +475,12 @@ async def get_swaps_summary( Summary statistics including volume, fees, success rate """ try: - async with db_manager.get_session_context() as session: - swap_repo = GatewaySwapRepository(session) - summary = await swap_repo.get_swaps_summary( - network=network, - wallet_address=wallet_address, - start_time=start_time, - end_time=end_time - ) - return summary + return await swap_service.get_swaps_summary( + network=network, + wallet_address=wallet_address, + start_time=start_time, + end_time=end_time + ) except Exception as e: logger.error(f"Error getting swaps summary: {e}", exc_info=True) diff --git a/routers/performance.py b/routers/performance.py new file mode 100644 index 00000000..e1cf3b18 --- /dev/null +++ b/routers/performance.py @@ -0,0 +1,265 @@ +"""Performance Router - one performance surface for both controllers and executors. + +The whole point of these routes is that a consumer writes seriesFor(scope) and +latestFor(scope) once each. The two populations live in different tables with different +shapes, and the branch between them is a query parameter rather than a path so it stays a +parameter to a client too. + +Two routes, mirroring the pair this replaces: `/history` for the series and `/latest` for +the current value of every scope. Together they are a complete substitute for +/bot-orchestration/controller-performance-history and -latest, so a consumer can move off +those entirely rather than straddling both surfaces. + +The controller subject of each goes through the existing BotsOrchestrator method the old +route already calls -- get_controller_performance_history and +get_latest_controller_performance -- so old and new share one query path by construction +and cannot drift. Those two existing routes are untouched and stay wire-compatible; this +is new surface only. +""" +import logging +from datetime import datetime, timezone +from typing import Optional + +from fastapi import APIRouter, Depends, HTTPException, Query + +from deps import get_bots_orchestrator, get_executor_service +from models.performance import ( + SUBJECT_CONTROLLER, + PerformanceHistoryResponse, + PerformanceLatestResponse, + controller_row_to_performance_row, + executor_row_to_performance_row, +) +from services.bots_orchestrator import BotsOrchestrator +from services.executor_service import ExecutorService + +logger = logging.getLogger(__name__) + +router = APIRouter(tags=["Performance"], prefix="/performance") + +# Which filters belong to which population. A controller_id means different things on +# either side -- an MQTT bot's controller and an in-process executor's controller tag are +# not guaranteed to name the same thing -- so it is a filter WITHIN a subject and never a +# key to join the two on. +_CONTROLLER_ONLY_FILTERS = ("bot_name",) +_EXECUTOR_ONLY_FILTERS = ("executor_id", "executor_type", "account_name", "connector_name", "trading_pair") + + +def _reject_foreign_filters(subject: str, **supplied) -> None: + """400 when a filter belongs to the other population. + + FastAPI cannot express "executor_id is only legal when subject=executor", so the + cross-check is explicit: a filter aimed at the wrong population would otherwise be + accepted and silently ignored, which reads as an empty result rather than a mistake. + Both routes share this one rule so they cannot drift into disagreeing about which + filter belongs where. + """ + wrong_subject = ( + _EXECUTOR_ONLY_FILTERS if subject == SUBJECT_CONTROLLER else _CONTROLLER_ONLY_FILTERS + ) + offending = [name for name in wrong_subject if supplied.get(name) is not None] + if offending: + raise HTTPException( + status_code=400, + detail=f"{', '.join(offending)} {'is' if len(offending) == 1 else 'are'} " + f"not a valid filter for subject={subject}", + ) + + +def _timestamp_key(row: dict) -> datetime: + """Sort key for the controller rows, whose timestamps arrive as ISO strings. + + An unparseable timestamp sorts last rather than raising: one malformed row must not + take down a dashboard's whole tile set. + """ + try: + return datetime.fromisoformat(str(row.get("timestamp")).replace("Z", "+00:00")) + except (TypeError, ValueError): + return datetime.min.replace(tzinfo=timezone.utc) + + +@router.get("/history", response_model=PerformanceHistoryResponse) +async def get_performance_history( + subject: str = Query(description='Which population to read: "controller" or "executor"', + pattern="^(controller|executor)$"), + bot_name: Optional[str] = Query(default=None, description="Filter by bot name (controller subject only)"), + controller_id: Optional[str] = Query(default=None, description="Filter by controller ID (either subject)"), + executor_id: Optional[str] = Query(default=None, description="Filter by executor ID (executor subject only)"), + executor_type: Optional[str] = Query(default=None, description="Filter by executor type (executor subject only)"), + account_name: Optional[str] = Query(default=None, description="Filter by account name (executor subject only)"), + connector_name: Optional[str] = Query(default=None, description="Filter by connector (executor subject only)"), + trading_pair: Optional[str] = Query(default=None, description="Filter by trading pair (executor subject only)"), + start_time: Optional[str] = Query(default=None, description="ISO 8601 start of the window"), + end_time: Optional[str] = Query(default=None, description="ISO 8601 end of the window"), + interval: str = Query(default="5m", pattern="^(1m|5m|15m|30m|1h|4h|12h|1d)$"), + limit: int = Query(default=100, le=1000), + cursor: Optional[str] = Query(default=None, description="Cursor from a previous page's next_cursor"), + bots_manager: BotsOrchestrator = Depends(get_bots_orchestrator), + executor_service: ExecutorService = Depends(get_executor_service), +): + """Historical performance for one subject, newest first, in one normalized row shape. + + Both subjects page identically: descending timestamp, `cursor` is the last row's + timestamp, `has_more` says whether another page exists. + + `interval` is a floor, not a guarantee. The controller series is written on a + 5-minute grain, so asking it for `1m` returns that native grain; the executor series + is written on PERFORMANCE_EXECUTOR_SNAPSHOT_INTERVAL (60s by default). The echoed + `interval` says what was asked for, the timestamps say what was served. + + It is a claim about resolution only. Sampling is per scope -- per executor, per (bot, + controller) -- so a coarser interval thins every scope's own series and never drops a + scope from the answer. An unnarrowed query over a fleet returns the whole fleet at + every interval, at a lower resolution. + + Realized vs unrealized, for executors: an executor reports one net PnL, and the split + comes from whether its position is really settled. A `POSITION_HOLD` close is NOT + settled -- it hands the position on to position_holds, which report it in their own + right -- so its PnL stays in `unrealized_pnl_quote` even though `is_terminal` is true + and `status` is TERMINATED. Counting it realized here would double-count it, the same + exclusion /executors/performance applies. Any other terminal close is realized. + + An executor's series is answered from executor_performance_snapshots alone, including + its final value: completion writes a terminal row, so there is no join to the + executors table and no "and then append the last point" rule. An executor that was + live when the API crashed gets its terminal row from the startup reap instead, carrying + `close_type: SYSTEM_CLEANUP` and the last figures observed before the crash -- an + approximated close, marked as one, rather than a series that never ends. + """ + _reject_foreign_filters( + subject, + bot_name=bot_name, + executor_id=executor_id, + executor_type=executor_type, + account_name=account_name, + connector_name=connector_name, + trading_pair=trading_pair, + ) + + try: + parsed_start = datetime.fromisoformat(start_time) if start_time else None + parsed_end = datetime.fromisoformat(end_time) if end_time else None + except ValueError as e: + raise HTTPException(status_code=400, detail=f"Invalid datetime format: {e}") + + try: + if subject == SUBJECT_CONTROLLER: + history, next_cursor, has_more = await bots_manager.get_controller_performance_history( + bot_name=bot_name, + controller_id=controller_id, + limit=limit, + cursor=cursor, + start_time=parsed_start, + end_time=parsed_end, + interval=interval, + ) + rows = [controller_row_to_performance_row(row) for row in history] + else: + history, next_cursor, has_more = await executor_service.get_executor_performance_history( + executor_id=executor_id, + executor_type=executor_type, + controller_id=controller_id, + account_name=account_name, + connector_name=connector_name, + trading_pair=trading_pair, + limit=limit, + cursor=cursor, + start_time=parsed_start, + end_time=parsed_end, + interval=interval, + ) + rows = [executor_row_to_performance_row(row) for row in history] + except Exception as e: + logger.error(f"Failed to get {subject} performance history: {e}", exc_info=True) + raise HTTPException(status_code=500, detail=str(e)) + + return PerformanceHistoryResponse( + status="success", + data=rows, + pagination={ + "next_cursor": next_cursor, + "has_more": has_more, + "limit": limit, + "interval": interval, + }, + ) + + +@router.get("/latest", response_model=PerformanceLatestResponse) +async def get_latest_performance( + subject: str = Query(description='Which population to read: "controller" or "executor"', + pattern="^(controller|executor)$"), + bot_name: Optional[str] = Query(default=None, description="Filter by bot name (controller subject only)"), + controller_id: Optional[str] = Query(default=None, description="Filter by controller ID (either subject)"), + executor_id: Optional[str] = Query(default=None, description="Filter by executor ID (executor subject only)"), + executor_type: Optional[str] = Query(default=None, description="Filter by executor type (executor subject only)"), + account_name: Optional[str] = Query(default=None, description="Filter by account name (executor subject only)"), + connector_name: Optional[str] = Query(default=None, description="Filter by connector (executor subject only)"), + trading_pair: Optional[str] = Query(default=None, description="Filter by trading pair (executor subject only)"), + limit: int = Query(default=100, le=1000, description="Cap on how many scopes come back, newest first"), + bots_manager: BotsOrchestrator = Depends(get_bots_orchestrator), + executor_service: ExecutorService = Depends(get_executor_service), +): + """The most recent snapshot of every scope in one subject, newest first. + + One row per scope -- per (bot, controller), or per executor -- in the same normalized + shape `/history` returns, so a dashboard's live tiles and its charts read the same + fields off the same client. + + `limit` caps how many scopes come back, it is not a page boundary: there is no cursor + here because this is not a series. Newest-first ordering means the scopes that are + still reporting come first, which matters far more for executors than for controllers + -- every executor that ever ran leaves a terminal row behind, so the executor + population grows without bound while the controller one does not. + + This reads the stored series, not live memory. An executor younger than one snapshot + interval has no row yet and does not appear, and a live one's figures are up to one + interval stale -- by design, so that this row and the last row of `/history` are the + same row. Live in-memory figures are what `/executors/` serves. + + A closed executor's latest row is its terminal row, carrying `is_terminal: true` and + its `close_type`, so "the final value" needs no separate call. That holds for every + way an executor can close, including the startup reap of one a crash left behind + (`close_type: SYSTEM_CLEANUP`): this row and `GET /executors/{executor_id}` are written + in the same transaction, so the two surfaces cannot disagree about whether an executor + is done. + """ + _reject_foreign_filters( + subject, + bot_name=bot_name, + executor_id=executor_id, + executor_type=executor_type, + account_name=account_name, + connector_name=connector_name, + trading_pair=trading_pair, + ) + + try: + if subject == SUBJECT_CONTROLLER: + # get_latest_controller_performance only narrows by bot_name -- it is the + # method the old route calls and is deliberately not changed. controller_id + # and the limit are applied here instead. That is cheap: the result is one row + # per (bot, controller), which is small by construction. + snapshots = await bots_manager.get_latest_controller_performance(bot_name=bot_name) + if controller_id: + snapshots = [s for s in snapshots if s.get("controller_id") == controller_id] + # sorted(), not .sort(): the list belongs to the orchestrator's caller and + # this route has no business reordering it in place. + ordered = sorted(snapshots, key=_timestamp_key, reverse=True) + rows = [controller_row_to_performance_row(row) for row in ordered[:limit]] + else: + snapshots = await executor_service.get_latest_executor_performance( + executor_id=executor_id, + executor_type=executor_type, + controller_id=controller_id, + account_name=account_name, + connector_name=connector_name, + trading_pair=trading_pair, + limit=limit, + ) + rows = [executor_row_to_performance_row(row) for row in snapshots] + except Exception as e: + logger.error(f"Failed to get latest {subject} performance: {e}", exc_info=True) + raise HTTPException(status_code=500, detail=str(e)) + + return PerformanceLatestResponse(status="success", data=rows) diff --git a/routers/websocket.py b/routers/websocket.py index 96d8d83d..6d36511a 100644 --- a/routers/websocket.py +++ b/routers/websocket.py @@ -7,8 +7,10 @@ import secrets import time import uuid +from typing import Optional, Tuple from fastapi import APIRouter, WebSocket, WebSocketDisconnect +from starlette.websockets import WebSocketState from config import settings from services.websocket_manager import WebSocketManager @@ -19,42 +21,106 @@ HEARTBEAT_INTERVAL = 30 # seconds +# Subprotocol a browser client offers to carry its credentials, since the JS WebSocket API +# cannot set an Authorization header: new WebSocket(url, ["hummingbot-auth", b64url(user:pass)]) +AUTH_SUBPROTOCOL = "hummingbot-auth" -def _authenticate_websocket(websocket: WebSocket) -> bool: + +def _decode_basic_credentials(encoded: str) -> Optional[Tuple[str, str]]: + """ + Decode a ``base64(username:password)`` blob into its two halves. + + Accepts base64url as well as standard base64, with or without padding: base64url is the + only variant whose alphabet is a valid ``Sec-WebSocket-Protocol`` token, so browser + clients have to use it, while ``Authorization: Basic`` uses the standard alphabet. + + Returns None if the blob is not decodable or carries no ``:`` separator. """ - Authenticate a WebSocket connection using Basic Auth from headers or query params. + padded = encoded + "=" * (-len(encoded) % 4) + try: + decoded = base64.b64decode(padded.replace("-", "+").replace("_", "/")).decode("utf-8") + except Exception: + return None + if ":" not in decoded: + return None + username, password = decoded.split(":", 1) + return username, password - Returns True if authenticated, False otherwise. + +def _authenticate_websocket(websocket: WebSocket) -> Tuple[bool, Optional[str]]: + """ + Authenticate a WebSocket handshake from its headers, before it is accepted. + + Credentials are never read from the query string: uvicorn logs the full path *with* its + query string for every handshake, so a ``?username=``/``?password=``/``?token=`` channel + would write the global admin credentials into access logs, traces and browser history. + + Two credential channels are supported, both header-based: + - ``Authorization: Basic base64(username:password)`` — same channel as the HTTP routes, + for any client that can set request headers. + - ``Sec-WebSocket-Protocol: hummingbot-auth, base64url(username:password)`` — for + browsers, whose ``new WebSocket(url, protocols)`` API cannot set headers but can + offer subprotocols. + + Returns ``(authenticated, subprotocol)``. When credentials arrived over the subprotocol + channel the selected subprotocol must be echoed back in ``websocket.accept()``, otherwise + the browser fails the connection itself. """ - # Try Authorization header first + credentials: Optional[Tuple[str, str]] = None + subprotocol: Optional[str] = None + auth_header = websocket.headers.get("authorization", "") if auth_header.startswith("Basic "): - try: - decoded = base64.b64decode(auth_header[6:]).decode("utf-8") - ws_user, ws_pass = decoded.split(":", 1) - except Exception: - return False + credentials = _decode_basic_credentials(auth_header[6:]) else: - # Fallback: ?token=base64(user:pass) query param - token = websocket.query_params.get("token") - if token: - try: - decoded = base64.b64decode(token).decode("utf-8") - ws_user, ws_pass = decoded.split(":", 1) - except Exception: - return False - else: - # Fallback to query parameters - ws_user = websocket.query_params.get("username", "") - ws_pass = websocket.query_params.get("password", "") - + offered = [ + protocol.strip() + for protocol in websocket.headers.get("sec-websocket-protocol", "").split(",") + if protocol.strip() + ] + if len(offered) >= 2 and offered[0] == AUTH_SUBPROTOCOL: + credentials = _decode_basic_credentials(offered[1]) + subprotocol = AUTH_SUBPROTOCOL + + if credentials is None: + return False, None + + ws_user, ws_pass = credentials correct_user = secrets.compare_digest( ws_user.encode(), settings.security.username.encode() ) correct_pass = secrets.compare_digest( ws_pass.encode(), settings.security.password.encode() ) - return correct_user and correct_pass + return bool(correct_user and correct_pass), subprotocol + + +async def _reject_unauthenticated(websocket: WebSocket) -> None: + """ + Refuse the handshake itself instead of accepting it and closing with 4001. + + An unauthenticated peer never reaches an open WebSocket: it gets an HTTP 401 handshake + response where the server supports the ASGI websocket denial-response extension, and a + 1008 policy-violation close (which the server turns into an HTTP 403) where it does not. + """ + if websocket.client_state == WebSocketState.CONNECTING: + await websocket.receive() + + if "websocket.http.response" in (websocket.scope.get("extensions") or {}): + await websocket.send({ + "type": "websocket.http.response.start", + "status": 401, + "headers": [ + (b"www-authenticate", b'Basic realm="hummingbot-api"'), + (b"content-type", b"text/plain; charset=utf-8"), + ], + }) + await websocket.send({ + "type": "websocket.http.response.body", + "body": b"Authentication failed", + }) + else: + await websocket.close(code=1008, reason="Authentication failed") async def _heartbeat_loop(websocket: WebSocket) -> None: @@ -77,8 +143,11 @@ async def market_data_websocket(websocket: WebSocket) -> None: """ WebSocket endpoint for streaming market data. - Authentication: Basic Auth via Authorization header, ?token=base64(user:pass), - or query params (?username=...&password=...). + Authentication (headers only, never the query string): + - Authorization: Basic base64(username:password) + - Sec-WebSocket-Protocol: hummingbot-auth, base64url(username:password) + for browsers; the server echoes back the "hummingbot-auth" subprotocol. + Unauthenticated handshakes are refused with HTTP 401, not accepted. Subscribe/unsubscribe protocol: -> {"action": "subscribe", "type": "candles", "connector": "binance", @@ -93,16 +162,13 @@ async def market_data_websocket(websocket: WebSocket) -> None: - order_book: order book snapshots with configurable depth - trades: real-time trade events """ - await websocket.accept() - - if not _authenticate_websocket(websocket): - await websocket.send_json({ - "type": "error", - "message": "Authentication failed", - }) - await websocket.close(code=4001, reason="Authentication failed") + authenticated, subprotocol = _authenticate_websocket(websocket) + if not authenticated: + await _reject_unauthenticated(websocket) return + await websocket.accept(subprotocol=subprotocol) + manager: WebSocketManager = websocket.app.state.websocket_manager conn_id = manager.generate_connection_id() @@ -158,8 +224,11 @@ async def executors_websocket(websocket: WebSocket) -> None: """ WebSocket endpoint for streaming executor data. - Authentication: Basic Auth via Authorization header, ?token=base64(user:pass), - or query params (?username=...&password=...). + Authentication (headers only, never the query string): + - Authorization: Basic base64(username:password) + - Sec-WebSocket-Protocol: hummingbot-auth, base64url(username:password) + for browsers; the server echoes back the "hummingbot-auth" subprotocol. + Unauthenticated handshakes are refused with HTTP 401, not accepted. Subscribe/unsubscribe protocol: -> {"action": "subscribe", "type": "executor_summary", "update_interval": 2.0} @@ -178,17 +247,13 @@ async def executors_websocket(websocket: WebSocket) -> None: - bot_status: single bot status with performance & custom_info (requires bot_name) - all_bots_status: all active bots status with performance & custom_info """ - await websocket.accept() - - # Authenticate - if not _authenticate_websocket(websocket): - await websocket.send_json({ - "type": "error", - "message": "Authentication failed", - }) - await websocket.close(code=4001, reason="Authentication failed") + authenticated, subprotocol = _authenticate_websocket(websocket) + if not authenticated: + await _reject_unauthenticated(websocket) return + await websocket.accept(subprotocol=subprotocol) + # Get manager from app state manager = websocket.app.state.executor_ws_manager conn_id = str(uuid.uuid4())[:12] diff --git a/scripts/backfill_gas_tokens.py b/scripts/backfill_gas_tokens.py new file mode 100644 index 00000000..3e796c89 --- /dev/null +++ b/scripts/backfill_gas_tokens.py @@ -0,0 +1,67 @@ +"""One-shot repair of gas_token on CLMM liquidity events (CORR-104). + +Rows written before ARCH-054 unified the chain -> native-gas-token map carry a wrong +gas_token: the add/remove handlers' old ternary left it NULL off solana/ethereum (and a +row inserted CONFIRMED is never re-polled, so the NULL is permanent), and the transaction +poller's old 6-entry dict wrote "UNKNOWN" for base, arbitrum and polygon. The write paths +are already correct; this script repairs the rows they left behind. + +The repair itself lives in GatewayCLMMRepository.backfill_liquidity_gas_tokens() so it is +testable; this module is only the entry point that opens a session and reports. + +Safe to re-run: a repaired row no longer matches the filter, so a second run is a no-op. +Rows whose chain still resolves to "UNKNOWN" are left untouched and reported, so an +unmapped chain surfaces instead of being papered over. + +Usage: + conda run --no-capture-output -n hummingbot-api python -m scripts.backfill_gas_tokens --dry-run + conda run --no-capture-output -n hummingbot-api python -m scripts.backfill_gas_tokens +""" +import argparse +import asyncio +import logging +import sys + +from config import settings +from database.connection import AsyncDatabaseManager +from database.repositories.gateway_clmm_repository import GatewayCLMMRepository + +logger = logging.getLogger("backfill_gas_tokens") + + +async def backfill(dry_run: bool = False) -> dict: + """Run the repair against the configured database, committing unless dry_run.""" + db_manager = AsyncDatabaseManager(settings.database.url) + try: + async with db_manager.get_session_context() as session: + repo = GatewayCLMMRepository(session) + report = await repo.backfill_liquidity_gas_tokens() + if dry_run: + # Discard the pending UPDATEs; the session context would commit them. + await session.rollback() + return report + finally: + await db_manager.close() + + +def main(argv=None) -> int: + parser = argparse.ArgumentParser(description=__doc__.splitlines()[0]) + parser.add_argument("--dry-run", action="store_true", + help="report what would change without writing") + args = parser.parse_args(argv) + + logging.basicConfig(level=logging.INFO, format="%(levelname)s %(message)s") + report = asyncio.run(backfill(dry_run=args.dry_run)) + + prefix = "[dry run] " if args.dry_run else "" + logger.info("%sgas_token repaired on %d liquidity event(s)", prefix, report["fixed"]) + if report["unresolved"]: + logger.warning( + "%s%d event(s) left untouched: no native gas token is mapped for %s", + prefix, report["unresolved"], ", ".join(report["unresolved_networks"]), + ) + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/services/backtesting_service.py b/services/backtesting_service.py index 606da1bc..c663996e 100644 --- a/services/backtesting_service.py +++ b/services/backtesting_service.py @@ -13,11 +13,26 @@ With the bulk on disk a resident task costs a few KB, so retention is a count of results rather than a memory budget, the full payload is rehydrated on read, and results survive a restart instead of dying with the process. + +The simulation itself runs in a spawned worker process, never on the API's event loop. +`run_backtesting` is awaitable but, once the candles are downloaded, it is an uninterrupted +CPU loop over every candle with no suspension point in it. Awaited inline it pinned the +loop thread at 100% for as long as the run lasted -- every endpoint, `/docs` included, +stopped answering -- and `DELETE /backtesting/tasks/{id}` was powerless, because asyncio +delivers a cancellation at an await and there was none to deliver it at. A worker process +fixes both: the loop only ever waits on a short poll, so cancellation lands within +milliseconds and the child can be signalled dead, which is the only thing that stops a +wedged CPU loop. It also caps how many runs are in flight and abandons one that overruns +its wall-clock budget. """ import asyncio import gzip import json import logging +import multiprocessing +import pickle +import time +import traceback import uuid from collections import OrderedDict from datetime import datetime, timezone @@ -25,9 +40,8 @@ from pathlib import Path from typing import Any, Dict, Optional -from hummingbot.strategy_v2.backtesting.backtesting_engine_base import BacktestingEngineBase - from config import settings +from services.candles_cache import CandlesCache, cache_key logger = logging.getLogger(__name__) @@ -35,6 +49,17 @@ # definite "it existed and is gone" instead of a bare 404 it cannot distinguish from a typo. _REAPED_MEMORY = 1000 +# How often the loop checks on the worker. It is the upper bound on how long a cancellation +# or a timeout takes to be acted on, and 20 wakeups a second cost nothing next to a run that +# lasts minutes. +_POLL_INTERVAL = 0.05 + +# How long a signalled worker is given to die before it is killed outright. +_TERM_GRACE = 1.0 + +# Prefix for the file a worker leaves its outcome in, inside the results directory. +_WORKER_PREFIX = ".worker-" + class BacktestTaskStatus(str, Enum): PENDING = "pending" @@ -61,6 +86,180 @@ def _json_default(obj: Any) -> Any: return str(obj) +class BacktestTimeout(Exception): + """A run exceeded its wall-clock budget and its worker was terminated.""" + + +# -- worker process -- +# +# Everything below runs in the child. The engine is built there and dies with it, so each +# run owns its BacktestingEngineBase outright: the shared-instance race CORR-060 closed with +# a lock cannot occur across processes at all. What that isolation used to cost was the +# per-instance BacktestingDataProvider.candles_feeds cache, which made every run re-download +# its full candle history; the download now goes through a process-external store of the +# candle data itself (services/candles_cache.py), so a sweep over one market downloads once. +# The store holds frames, never an engine or a provider, and hands each reader its own copy +# -- so nothing mutable is reachable from two runs and CORR-060's guarantee is untouched. + + +def _shape_result(backtesting_results: dict) -> dict: + """Turn the engine's raw output into the JSON-ready payload the API returns.""" + processed_data = backtesting_results["processed_data"]["features"].fillna(0) + executors_info = [e.to_dict() for e in backtesting_results["executors"]] + results = backtesting_results["results"] + results["sharpe_ratio"] = results["sharpe_ratio"] if results["sharpe_ratio"] is not None else 0 + + # Serialize position holds + position_holds = [] + for ph in backtesting_results.get("position_holds", []): + position_holds.append({ + "connector_name": ph.connector_name, + "trading_pair": ph.trading_pair, + "buy_amount_base": float(ph.buy_amount_base), + "buy_amount_quote": float(ph.buy_amount_quote), + "sell_amount_base": float(ph.sell_amount_base), + "sell_amount_quote": float(ph.sell_amount_quote), + "net_amount_base": float(ph.net_amount_base), + "cum_fees_quote": float(ph.cum_fees_quote), + "volume_traded_quote": float(ph.volume_traded_quote), + "is_closed": ph.is_closed, + "n_executors": len(ph.source_executor_ids), + }) + + return { + "executors": executors_info, + "processed_data": processed_data.to_dict(), + "results": results, + "position_holds": position_holds, + "position_held_timeseries": backtesting_results.get("position_held_timeseries", []), + "pnl_timeseries": backtesting_results.get("pnl_timeseries", []), + } + + +def _default_candles_cache() -> CandlesCache: + return CandlesCache( + path=settings.backtesting.candles_cache_path, + max_entries=settings.backtesting.candles_cache_entries, + ttl_seconds=settings.backtesting.candles_cache_ttl_seconds, + ) + + +def _install_candle_cache(provider, cache: CandlesCache) -> None: + """Route this run's candle downloads through the shared store. + + The engine builds its own BacktestingDataProvider and hands it to the controller, so + wrapping the instance's get_candles_feed catches every download the run makes -- the + backtesting-resolution feed and each of the controller's own. + + The key is the whole of what the download depends on: the market, the interval, the + max_records that sets how far before the window the fetch reaches back, and the run's + window itself. An entry is therefore only ever served to a request for exactly the range + it holds; a wider or shifted window misses and is fetched. + + `loaded` is this run's own memo, so a controller asking twice for the same feed does not + re-read the file. It is a local, lives as long as the call chain that made it, and is + never seen by another run -- the cache deliberately keeps no cross-run object graph. + """ + download = provider.get_candles_feed + loaded = {} + + async def get_candles_feed(config): + key = cache_key( + config.connector, config.trading_pair, config.interval, config.max_records, + provider.start_time, provider.end_time, + ) + frame = loaded.get(key) + if frame is None: + frame = cache.get(key) + if frame is None: + # Whoever waits here finds the entry the first one wrote, so N workers that miss + # together still download once -- the check-then-fetch race CORR-061 closed for + # live feeds, in its offline form. + with cache.single_flight(key): + frame = cache.get(key) + if frame is None: + feed_key = provider._generate_candle_feed_key(config) + held = provider.candles_feeds.get(feed_key) + frame = await download(config) + # Upstream answers from its own in-run dict whenever a feed it already + # holds spans the window, whatever max_records was asked for. Such a + # frame was fetched for a different request, so storing it here would + # let a later run be served less history than it asked for. Only what + # was genuinely downloaded is worth keeping. + if frame is not held: + cache.put(key, frame) + loaded[key] = frame + provider.candles_feeds[provider._generate_candle_feed_key(config)] = frame + return frame + + provider.get_candles_feed = get_candles_feed + + +def _run_backtest_blocking(config: dict, controllers_path: str, controllers_module: str, + cache: Optional[CandlesCache] = None) -> dict: + """Build the controller config, run the simulation to completion, shape the payload.""" + # Imported here rather than at module scope so only a process that actually backtests + # pays for pulling in hummingbot. + from hummingbot.strategy_v2.backtesting.backtesting_engine_base import BacktestingEngineBase + + engine = BacktestingEngineBase() + _install_candle_cache(engine.backtesting_data_provider, cache if cache is not None else _default_candles_cache()) + if isinstance(config["config"], str): + controller_config = engine.get_controller_config_instance_from_yml( + config_path=config["config"], + controllers_conf_dir_path=controllers_path, + controllers_module=controllers_module + ) + else: + controller_config = engine.get_controller_config_instance_from_dict( + config_data=config["config"], + controllers_module=controllers_module + ) + backtesting_results = asyncio.run(engine.run_backtesting( + controller_config=controller_config, + trade_cost=config.get("trade_cost", 0.0006), + start=int(config["start_time"]), + end=int(config["end_time"]), + backtesting_resolution=config.get("backtesting_resolution", "1m"), + )) + return _shape_result(backtesting_results) + + +def _worker_main(config: dict, controllers_path: str, controllers_module: str, out_path: str) -> None: + """Worker entry point: run one backtest and leave a pickled envelope at out_path. + + The outcome goes through a file rather than a pipe because a result is megabytes and a + pipe holds only tens of kilobytes: a child blocking in send() against a parent that is + waiting for it to exit is a deadlock, and the parent here waits on the process, not on + a read. The file is written whole before the child exits, so its absence afterwards + means the worker died rather than finished. + """ + try: + result = _run_backtest_blocking(config, controllers_path, controllers_module) + blob = pickle.dumps({"ok": True, "result": result}, protocol=pickle.HIGHEST_PROTOCOL) + except Exception as e: + blob = pickle.dumps({"ok": False, "error": f"{type(e).__name__}: {e}", "traceback": traceback.format_exc()}) + with open(out_path, "wb") as fh: + fh.write(blob) + + +def _terminate(proc: multiprocessing.Process) -> None: + """Stop a worker and reap it. A CPU-bound child is signalled, not asked.""" + if proc.pid is None: + return + try: + if proc.is_alive(): + proc.terminate() + proc.join(_TERM_GRACE) + if proc.is_alive(): + proc.kill() + proc.join(_TERM_GRACE) + proc.join(0) + proc.close() + except (OSError, ValueError) as e: # already reaped, or closed under us + logger.debug(f"Backtest worker cleanup: {e}") + + class BacktestTask: def __init__(self, task_id: str, config: dict): self.task_id = task_id @@ -95,13 +294,32 @@ def metadata(self) -> dict: class BacktestingService: - def __init__(self, max_results: Optional[int] = None, results_path: Optional[str] = None): + def __init__( + self, + max_results: Optional[int] = None, + results_path: Optional[str] = None, + max_concurrent: Optional[int] = None, + timeout_seconds: Optional[float] = None, + ): self._tasks: "OrderedDict[str, BacktestTask]" = OrderedDict() - self._engine = BacktestingEngineBase() + # Each run gets its own engine, in its own process. An engine is mutated in place by + # run_backtesting -- the time window, the controller, the resolution and the per-run + # accumulators all live on self -- and then suspends for seconds on the candle + # download, so sharing one instance between overlapping runs returned silently wrong + # numbers (CORR-060). A process boundary makes that unshareable by construction. + self._worker_target = _worker_main + # Runs beyond the cap queue on this rather than piling a core each onto the box. + max_concurrent = max_concurrent if max_concurrent is not None else settings.backtesting.max_concurrent + self._max_concurrent = max(1, int(max_concurrent)) + self._slots = asyncio.Semaphore(self._max_concurrent) + self._timeout = ( + timeout_seconds if timeout_seconds is not None else settings.backtesting.timeout_seconds + ) self._max_results = max_results if max_results is not None else settings.backtesting.max_results self._results_dir = Path(results_path if results_path is not None else settings.backtesting.results_path) self._reaped: "OrderedDict[str, str]" = OrderedDict() self._results_dir.mkdir(parents=True, exist_ok=True) + self._clear_worker_files() self._restore() # Honour a limit that was lowered since the last run. self._reap() @@ -143,7 +361,14 @@ def was_reaped(self, task_id: str) -> bool: return task_id in self._reaped def cancel_task(self, task_id: str) -> bool: - """Cancel a running task or remove a completed one, discarding its archive.""" + """Cancel a running task or remove a completed one, discarding its archive. + + Cancelling the coroutine is enough to stop the computation now that the coroutine + actually suspends: the cancellation is delivered at the next poll, at most + _POLL_INTERVAL away, and _run_in_worker's finally signals the worker dead. Before the + run moved off the loop this call was a lie -- it reported CANCELLED while the + simulation kept the process pinned, because there was no await to deliver it at. + """ task = self._tasks.get(task_id) if task is None: return False @@ -189,54 +414,58 @@ async def _run_task(self, task: BacktestTask): async def _execute_backtest(self, config: dict) -> dict: """Core backtest execution logic shared by sync and async modes.""" - if isinstance(config["config"], str): - controller_config = self._engine.get_controller_config_instance_from_yml( - config_path=config["config"], - controllers_conf_dir_path=settings.app.controllers_path, - controllers_module=settings.app.controllers_module - ) - else: - controller_config = self._engine.get_controller_config_instance_from_dict( - config_data=config["config"], - controllers_module=settings.app.controllers_module - ) - backtesting_results = await self._engine.run_backtesting( - controller_config=controller_config, - trade_cost=config.get("trade_cost", 0.0006), - start=int(config["start_time"]), - end=int(config["end_time"]), - backtesting_resolution=config.get("backtesting_resolution", "1m"), + async with self._slots: + return await self._run_in_worker(config) + + async def _run_in_worker(self, config: dict) -> dict: + """Run one backtest in a child process, supervised from the loop. + + The loop never touches the simulation: it waits in short sleeps, which is what makes + the API stay responsive, makes a cancellation land within a poll interval, and lets + the wall-clock budget be enforced. Whatever ends the wait -- success, timeout, + cancellation, or the caller hanging up -- the worker is signalled dead on the way out, + so no orphan is left burning a core. + """ + ctx = multiprocessing.get_context("spawn") + out_path = self._results_dir / f"{_WORKER_PREFIX}{uuid.uuid4().hex}.pkl" + proc = ctx.Process( + target=self._worker_target, + args=(config, settings.app.controllers_path, settings.app.controllers_module, str(out_path)), + daemon=True, ) - processed_data = backtesting_results["processed_data"]["features"].fillna(0) - executors_info = [e.to_dict() for e in backtesting_results["executors"]] - results = backtesting_results["results"] - results["sharpe_ratio"] = results["sharpe_ratio"] if results["sharpe_ratio"] is not None else 0 - - # Serialize position holds - position_holds = [] - for ph in backtesting_results.get("position_holds", []): - position_holds.append({ - "connector_name": ph.connector_name, - "trading_pair": ph.trading_pair, - "buy_amount_base": float(ph.buy_amount_base), - "buy_amount_quote": float(ph.buy_amount_quote), - "sell_amount_base": float(ph.sell_amount_base), - "sell_amount_quote": float(ph.sell_amount_quote), - "net_amount_base": float(ph.net_amount_base), - "cum_fees_quote": float(ph.cum_fees_quote), - "volume_traded_quote": float(ph.volume_traded_quote), - "is_closed": ph.is_closed, - "n_executors": len(ph.source_executor_ids), - }) - - return { - "executors": executors_info, - "processed_data": processed_data.to_dict(), - "results": results, - "position_holds": position_holds, - "position_held_timeseries": backtesting_results.get("position_held_timeseries", []), - "pnl_timeseries": backtesting_results.get("pnl_timeseries", []), - } + proc.start() + deadline = time.monotonic() + self._timeout + try: + while proc.is_alive(): + if time.monotonic() >= deadline: + raise BacktestTimeout( + f"Backtest exceeded its wall-clock budget of {self._timeout:g}s and was terminated" + ) + await asyncio.sleep(_POLL_INTERVAL) + return self._read_outcome(out_path, proc.exitcode) + finally: + _terminate(proc) + try: + out_path.unlink(missing_ok=True) + except OSError as e: + logger.warning(f"Could not remove backtest worker file {out_path}: {e}") + + @staticmethod + def _read_outcome(out_path: Path, exitcode: Optional[int]) -> dict: + """Unwrap what the worker left behind, or explain why there is nothing to unwrap.""" + if not out_path.exists(): + raise RuntimeError( + f"Backtest worker exited with code {exitcode} without producing a result" + ) + try: + with open(out_path, "rb") as fh: + envelope = pickle.load(fh) + except (OSError, pickle.UnpicklingError, EOFError, AttributeError) as e: + raise RuntimeError(f"Could not read the backtest worker result: {e}") + if not envelope.get("ok"): + logger.error(f"Backtest worker failed:\n{envelope.get('traceback', '')}") + raise RuntimeError(envelope.get("error", "backtest failed in the worker process")) + return envelope["result"] # -- archive -- @@ -280,6 +509,14 @@ def _persist_index(self) -> None: except (OSError, TypeError, ValueError) as e: logger.error(f"Could not persist backtest index: {e}") + def _clear_worker_files(self) -> None: + """Drop worker outcome files a previous process died before collecting.""" + for path in self._results_dir.glob(f"{_WORKER_PREFIX}*.pkl"): + try: + path.unlink() + except OSError as e: + logger.warning(f"Could not remove stale backtest worker file {path}: {e}") + def _restore(self) -> None: """Rebuild finished tasks from the index so results outlive a restart.""" if not self._index_path.exists(): diff --git a/services/bots_orchestrator.py b/services/bots_orchestrator.py index 6392701e..85c9fa3c 100644 --- a/services/bots_orchestrator.py +++ b/services/bots_orchestrator.py @@ -472,6 +472,20 @@ async def mark_bot_run_stopped(self, bot_name: str, final_status: Optional[Dict] await bot_run_repo.update_bot_run_stopped(bot_name, final_status=final_status) logger.info(f"Updated bot run status to STOPPED for {bot_name}") + async def mark_bot_run_errored(self, bot_name: str, error_message: str): + """Close out a bot run with an error state and message. + + Best-effort: a failure to record the error must never mask the original + one, so database problems are logged and swallowed. + """ + try: + async with self.db_manager.get_session_context() as session: + bot_run_repo = BotRunRepository(session) + await bot_run_repo.update_bot_run_stopped(bot_name, error_message=error_message) + logger.info(f"Updated bot run with error status for {bot_name}: {error_message}") + except Exception as e: + logger.error(f"Failed to update bot run with error: {e}") + async def get_bot_runs( self, bot_name: Optional[str] = None, @@ -521,6 +535,23 @@ async def get_bot_run_by_id(self, bot_run_id: int) -> Optional[Dict]: return None return self._serialize_bot_run(bot_run) + async def delete_bot_runs_for_bot(self, bot_name: str) -> int: + """Delete every bot run recorded under a bot name, returning how many went. + + Best-effort: the caller has already deleted the archived bot's files by the time + this runs, so a database problem must not report that deletion as a failure. + """ + try: + async with self.db_manager.get_session_context() as session: + bot_run_repo = BotRunRepository(session) + deleted = await bot_run_repo.delete_bot_runs_by_bot_name(bot_name) + if deleted > 0: + logger.info(f"Deleted {deleted} bot run record(s) for '{bot_name}'") + return deleted + except Exception as e: + logger.warning(f"Failed to clean bot run records for '{bot_name}': {e}") + return 0 + async def delete_bot_run(self, bot_run_id: int) -> Optional[Dict]: """Delete a bot run record and its archived folder. @@ -601,8 +632,6 @@ def _serialize_bot_run(run, include_final_status: bool = True) -> Dict: async def stop_and_archive_bot( self, bot_name: str, - container_name: str, - bot_name_for_orchestrator: str, skip_order_cancellation: bool, archive_locally: bool, s3_bucket: Optional[str], @@ -618,13 +647,13 @@ async def stop_and_archive_bot( logger.info(f"Starting background stop-and-archive for {bot_name}") # Step 1: Capture bot final status before stopping (while bot is still running) - logger.info(f"Capturing final status for {bot_name_for_orchestrator}") + logger.info(f"Capturing final status for {bot_name}") final_status = None try: - final_status = self.get_bot_status(bot_name_for_orchestrator) - logger.info(f"Captured final status for {bot_name_for_orchestrator}: {final_status}") + final_status = self.get_bot_status(bot_name) + logger.info(f"Captured final status for {bot_name}: {final_status}") except Exception as e: - logger.warning(f"Failed to capture final status for {bot_name_for_orchestrator}: {e}") + logger.warning(f"Failed to capture final status for {bot_name}: {e}") # Step 2: Update bot run with stopped_at timestamp and final status before stopping try: @@ -635,10 +664,10 @@ async def stop_and_archive_bot( # Continue with stop process even if database update fails # Step 3: Mark the bot as stopping, and stop the bot trading process - self.set_bot_stopping(bot_name_for_orchestrator) - logger.info(f"Stopping bot trading process for {bot_name_for_orchestrator}") + self.set_bot_stopping(bot_name) + logger.info(f"Stopping bot trading process for {bot_name}") stop_response = await self.stop_bot( - bot_name_for_orchestrator, + bot_name, skip_order_cancellation=skip_order_cancellation, async_backend=True # Always use async for background tasks ) @@ -646,6 +675,7 @@ async def stop_and_archive_bot( if not stop_response or not stop_response.get("success", False): error_msg = stop_response.get('error', 'Unknown error') if stop_response else 'No response from bot orchestrator' logger.error(f"Failed to stop bot process: {error_msg}") + await self.mark_bot_run_errored(bot_name, f"Failed to stop bot process: {error_msg}") return # Step 4: Wait for graceful shutdown (15 seconds as requested) @@ -658,44 +688,47 @@ async def stop_and_archive_bot( container_stopped = False for i in range(max_retries): - logger.info(f"Attempting to stop container {container_name} (attempt {i+1}/{max_retries})") - docker_manager.stop_container(container_name) + logger.info(f"Attempting to stop container {bot_name} (attempt {i+1}/{max_retries})") + docker_manager.stop_container(bot_name) # Check if container is already stopped - container_status = docker_manager.get_container_status(container_name) + container_status = docker_manager.get_container_status(bot_name) if container_status.get("state", {}).get("status") == "exited": container_stopped = True - logger.info(f"Container {container_name} is already stopped") + logger.info(f"Container {bot_name} is already stopped") break await asyncio.sleep(retry_interval) if not container_stopped: - logger.error(f"Failed to stop container {container_name} after {max_retries} attempts") + logger.error(f"Failed to stop container {bot_name} after {max_retries} attempts") + await self.mark_bot_run_errored( + bot_name, f"Failed to stop container after {max_retries} attempts" + ) return # Step 6: Archive the bot data - instance_dir = os.path.join('bots', 'instances', container_name) + instance_dir = os.path.join('bots', 'instances', bot_name) logger.info(f"Archiving bot data from {instance_dir}") try: if archive_locally: - bot_archiver.archive_locally(container_name, instance_dir) + bot_archiver.archive_locally(bot_name, instance_dir) else: - bot_archiver.archive_and_upload(container_name, instance_dir, bucket_name=s3_bucket) - logger.info(f"Successfully archived bot data for {container_name}") + bot_archiver.archive_and_upload(bot_name, instance_dir, bucket_name=s3_bucket) + logger.info(f"Successfully archived bot data for {bot_name}") except Exception as e: logger.error(f"Archive failed: {str(e)}") # Continue with removal even if archive fails # Step 7: Remove the container - logging.info(f"Removing container {container_name}") - remove_response = docker_manager.remove_container(container_name, force=False) + logging.info(f"Removing container {bot_name}") + remove_response = docker_manager.remove_container(bot_name, force=False) if not remove_response.get("success"): # If graceful remove fails, try force remove logging.warning("Graceful container removal failed, attempting force removal") - remove_response = docker_manager.remove_container(container_name, force=True) + remove_response = docker_manager.remove_container(bot_name, force=True) if remove_response.get("success"): logging.info(f"Successfully completed stop-and-archive for bot {bot_name}") @@ -709,41 +742,21 @@ async def stop_and_archive_bot( except Exception as e: logger.error(f"Failed to update bot run to archived: {e}") else: - logging.error(f"Failed to remove container {container_name}") + logging.error(f"Failed to remove container {bot_name}") - # Update bot run with error status (but keep stopped_at timestamp from earlier) - try: - async with self.db_manager.get_session_context() as session: - bot_run_repo = BotRunRepository(session) - await bot_run_repo.update_bot_run_stopped( - bot_name, - error_message="Failed to remove container during archive process" - ) - logger.info(f"Updated bot run with error status for {bot_name}") - except Exception as e: - logger.error(f"Failed to update bot run with error: {e}") + # Keep the stopped_at timestamp from earlier, but flag why it never archived + await self.mark_bot_run_errored(bot_name, "Failed to remove container during archive process") except Exception as e: logging.error(f"Error in background stop-and-archive for {bot_name}: {str(e)}") - - # Update bot run with error status - try: - async with self.db_manager.get_session_context() as session: - bot_run_repo = BotRunRepository(session) - await bot_run_repo.update_bot_run_stopped( - bot_name, - error_message=str(e) - ) - logger.info(f"Updated bot run with error status for {bot_name}") - except Exception as db_error: - logger.error(f"Failed to update bot run with error: {db_error}") + await self.mark_bot_run_errored(bot_name, str(e)) finally: # Always clear the stopping status when the background task completes - self.clear_bot_stopping(bot_name_for_orchestrator) + self.clear_bot_stopping(bot_name) logger.info(f"Cleared stopping status for bot {bot_name}") # Remove bot from active_bots and clear all MQTT data - if bot_name_for_orchestrator in self.active_bots: - self.mqtt_manager.clear_bot_data(bot_name_for_orchestrator) - del self.active_bots[bot_name_for_orchestrator] - logger.info(f"Removed bot {bot_name_for_orchestrator} from active_bots and cleared MQTT data") + if bot_name in self.active_bots: + self.mqtt_manager.clear_bot_data(bot_name) + del self.active_bots[bot_name] + logger.info(f"Removed bot {bot_name} from active_bots and cleared MQTT data") diff --git a/services/candles_cache.py b/services/candles_cache.py new file mode 100644 index 00000000..3a8e4b3f --- /dev/null +++ b/services/candles_cache.py @@ -0,0 +1,197 @@ +"""A bounded, process-external store for downloaded historical candles. + +Each backtest runs in its own worker process, which is what makes a run killable and +budgeted (ARCH-063) and what makes two runs structurally unable to corrupt each other +(CORR-060). The cost of that isolation is that `BacktestingDataProvider.candles_feeds` -- +the per-instance dict that used to let a second backtest over the same market skip the +download -- dies with the process. An optimizer sweeping N configs over one market +therefore downloaded the same candle history N times. + +This puts the cache back, outside the process, and stores the only thing that is safe to +share: the immutable *data*. Nothing here is an engine, a controller or a data provider, +and a reader gets its own unpickled copy of a frame -- so a run can do what it likes with +what it is handed, and no state is reachable from two runs at once. That is the property +CORR-060 established and this must not undo. + +Three rules make it safe to serve: + +- **Range-exact keys.** An entry is keyed by everything the download is a function of -- + connector, pair, interval, max_records and the run's window -- so a hit covers exactly + the range that was asked for. A window the cache does not hold is a miss, never a + narrower frame passed off as a wider one. +- **A freshness bound.** A window ending near "now" is downloaded with its last candle + still forming, so an entry is only served for `ttl_seconds` after it was fetched. Past + that it is a miss and the data is fetched again. +- **A size bound.** The number of entries is capped and the least recently used are + dropped, so sweeping across many pairs and timeframes cannot grow the store without + limit. + +Single-flight: the entry file doubles as the lock, so N workers that all miss the same key +at once produce one download rather than N. A caller whose filesystem has no `flock` still +gets correct results -- just, at worst, a duplicated first download. + +A cache is an optimization and never a failure mode: every read and write swallows its own +I/O errors and degrades to "download it again". +""" +import hashlib +import logging +import os +import pickle +import time +import uuid +from contextlib import contextmanager +from pathlib import Path +from typing import Any, Optional + +try: + import fcntl +except ImportError: # non-POSIX: single-flight degrades to "everyone downloads once" + fcntl = None + +logger = logging.getLogger(__name__) + +# How long a leftover half-written entry from a killed worker is left alone before it is +# swept. A write takes seconds, so anything older than this is debris, not a live write. +_TMP_GRACE = 600.0 + +_ENTRY_SUFFIX = ".pkl" +_TMP_SUFFIX = ".pkl.tmp" + + +def cache_key(*parts: Any) -> str: + """A stable key over everything a download is a function of.""" + joined = "|".join(str(part) for part in parts) + return hashlib.sha256(joined.encode("utf-8")).hexdigest()[:32] + + +class CandlesCache: + """Downloaded candle frames on disk, keyed by exact range, bounded in count and age.""" + + def __init__(self, path: str, max_entries: int, ttl_seconds: float): + self._path = Path(path) + self._max_entries = int(max_entries) + self._ttl = float(ttl_seconds) + if self.enabled: + try: + self._path.mkdir(parents=True, exist_ok=True) + except OSError as e: + logger.warning(f"Candle cache at {self._path} is unusable, running without it: {e}") + self._max_entries = 0 + + @property + def enabled(self) -> bool: + """A cap of zero turns the cache off entirely -- the escape hatch for an operator.""" + return self._max_entries > 0 + + def _entry_path(self, key: str) -> Path: + return self._path / f"{key}{_ENTRY_SUFFIX}" + + def get(self, key: str) -> Optional[Any]: + """The frame stored under this exact key, or None if it is absent, stale or unreadable. + + Never deletes: a caller that misses goes on to overwrite the entry anyway, and + removing a file another worker is holding open as its single-flight lock would only + cost a duplicate download. + """ + if not self.enabled: + return None + path = self._entry_path(key) + try: + with open(path, "rb") as fh: + envelope = pickle.load(fh) + except FileNotFoundError: + return None + except Exception as e: # truncated, half-written, or written by another pandas + logger.debug(f"Candle cache entry {key} is unreadable, refetching: {e}") + return None + if not isinstance(envelope, dict) or "frame" not in envelope: + return None + if time.time() - float(envelope.get("created_at", 0)) > self._ttl: + return None + # Recency for eviction is the file's mtime, kept apart from the fetch time inside + # the envelope so that reading a hot entry cannot extend its freshness. + try: + os.utime(path, None) + except OSError: + pass + return envelope["frame"] + + def put(self, key: str, frame: Any) -> None: + """Store a frame under an exact key, atomically, then honour the size bound.""" + if not self.enabled: + return + tmp = self._path / f"{key}.{os.getpid()}.{uuid.uuid4().hex}{_TMP_SUFFIX}" + try: + with open(tmp, "wb") as fh: + pickle.dump({"created_at": time.time(), "frame": frame}, fh, protocol=pickle.HIGHEST_PROTOCOL) + # Replace is atomic, so a concurrent reader sees either the old entry or the + # new one, never a partial write. + os.replace(tmp, self._entry_path(key)) + except Exception as e: + logger.warning(f"Could not cache candles under {key}: {e}") + try: + tmp.unlink(missing_ok=True) + except OSError: + pass + return + self._evict() + + @contextmanager + def single_flight(self, key: str): + """Serialize the workers that miss the same key, so only the first one downloads. + + The entry file is its own lock: taking it creates an empty placeholder, which reads + as a miss until it is replaced by a real entry, and which the size bound sweeps like + any other entry. Callers must re-check the cache inside this block -- the whole point + is that whoever waited here finds the entry the first one just wrote. + + The lock is advisory and best-effort: without `flock` support the block is a no-op + and the only consequence is that a first download happens more than once. + """ + if not self.enabled or fcntl is None: + yield + return + handle = None + try: + handle = open(self._entry_path(key), "a+b") + fcntl.flock(handle.fileno(), fcntl.LOCK_EX) + except OSError as e: + logger.debug(f"Candle cache could not lock {key}, downloading unguarded: {e}") + if handle is not None: + handle.close() + handle = None + try: + yield + finally: + if handle is not None: + try: + fcntl.flock(handle.fileno(), fcntl.LOCK_UN) + except OSError: + pass + handle.close() + + def _evict(self) -> None: + """Drop the least recently used entries beyond the cap, and any abandoned writes.""" + try: + entries = [(path.stat().st_mtime, path) for path in self._path.glob(f"*{_ENTRY_SUFFIX}")] + except OSError as e: + logger.debug(f"Could not scan the candle cache at {self._path}: {e}") + return + entries.sort() + for _, path in entries[: max(0, len(entries) - self._max_entries)]: + self._unlink(path) + + cutoff = time.time() - max(self._ttl, _TMP_GRACE) + for path in self._path.glob(f"*{_TMP_SUFFIX}"): + try: + if path.stat().st_mtime < cutoff: + self._unlink(path) + except OSError: + continue + + @staticmethod + def _unlink(path: Path) -> None: + try: + path.unlink(missing_ok=True) + except OSError as e: + logger.debug(f"Could not drop candle cache file {path}: {e}") diff --git a/services/docker_service.py b/services/docker_service.py index c6f346ca..41d400f2 100644 --- a/services/docker_service.py +++ b/services/docker_service.py @@ -6,11 +6,12 @@ from typing import Dict import docker -from docker.errors import DockerException +from docker.errors import DockerException, NotFound from docker.types import LogConfig from config import settings from models import V2ControllerDeployment +from models.bot_orchestration import validate_safe_config_name from utils.file_system import fs_util from utils.gateway_certs import ensure_gateway_certs, gateway_certs_dir @@ -37,6 +38,20 @@ def __init__(self): except DockerException as e: logger.error(f"It was not possible to connect to Docker. Please make sure Docker is running. Error: {e}") + @staticmethod + def _failure(message: str, error: str = "docker_error") -> Dict: + """A failure a route can turn into a status code, without the daemon's internals. + + docker-py's exception strings carry the socket URL and the negotiated API version + ("404 Client Error for http+docker://localhost/v1.55/containers/x/json: Not Found + ..."), so returning str(e) as the response body both published how this API talks + to its daemon and, because these routes returned it with a 200, let a caller that + checks the status code -- the normal way to detect failure -- read a container + that was never stopped as one that was. The raw error goes to the log; the caller + gets the kind of failure, which routers/docker.py maps to a status code. + """ + return {"success": False, "error": error, "message": message} + def get_active_containers(self, name_filter: str = None): try: all_containers = self.client.containers.list(filters={"status": "running"}) @@ -62,14 +77,16 @@ def get_active_containers(self, name_filter: str = None): ] return containers_info except DockerException as e: - return str(e) + logger.error(f"Error listing active containers: {e}") + return self._failure("Could not list running containers") def get_available_images(self): try: images = self.client.images.list() return {"images": images} except DockerException as e: - return str(e) + logger.error(f"Error listing images: {e}") + return self._failure("Could not list Docker images") def pull_image(self, image_name): try: @@ -110,13 +127,16 @@ def get_exited_containers(self, name_filter: str = None): ] return containers_info except DockerException as e: - return str(e) + logger.error(f"Error listing exited containers: {e}") + return self._failure("Could not list exited containers") def clean_exited_containers(self): try: self.client.containers.prune() + return {"success": True, "message": "Exited containers removed."} except DockerException as e: - return str(e) + logger.error(f"Error pruning exited containers: {e}") + return self._failure("Could not remove exited containers") def is_docker_running(self): try: @@ -126,18 +146,33 @@ def is_docker_running(self): return False def stop_container(self, container_name): + """Stop a running container. + + Reports failure in-band rather than raising: stop-and-archive calls this in a + retry loop and decides whether it worked by reading the container's status + afterwards, so an exception here would abort a workflow that is designed to + tolerate a stop that did not take. + """ try: container = self.client.containers.get(container_name) container.stop() + return {"success": True, "message": f"Container {container_name} stopped."} + except NotFound: + return self._failure(f"No such container: {container_name}", error="not_found") except DockerException as e: - return str(e) + logger.error(f"Error stopping container {container_name}: {e}") + return self._failure(f"Could not stop container '{container_name}'") def start_container(self, container_name): try: container = self.client.containers.get(container_name) container.start() + return {"success": True, "message": f"Container {container_name} started."} + except NotFound: + return self._failure(f"No such container: {container_name}", error="not_found") except DockerException as e: - return str(e) + logger.error(f"Error starting container {container_name}: {e}") + return self._failure(f"Could not start container '{container_name}'") def get_container_status(self, container_name): """Get the status of a container""" @@ -151,16 +186,22 @@ def get_container_status(self, container_name): "exit_code": getattr(container.attrs.get("State", {}), "ExitCode", None) } } + except NotFound: + return self._failure(f"No such container: {container_name}", error="not_found") except DockerException as e: - return {"success": False, "message": str(e)} + logger.error(f"Error reading status of container {container_name}: {e}") + return self._failure(f"Could not read the status of container '{container_name}'") def remove_container(self, container_name, force=True): try: container = self.client.containers.get(container_name) container.remove(force=force) return {"success": True, "message": f"Container {container_name} removed successfully."} + except NotFound: + return self._failure(f"No such container: {container_name}", error="not_found") except DockerException as e: - return {"success": False, "message": str(e)} + logger.error(f"Error removing container {container_name}: {e}") + return self._failure(f"Could not remove container '{container_name}'") @staticmethod def _ensure_contained(path: str, base_dir: str, label: str): @@ -174,6 +215,20 @@ def _ensure_contained(path: str, base_dir: str, label: str): raise ValueError(f"Invalid {label}: '{path}' resolves outside of '{base_dir}'.") return resolved_path + @classmethod + def resolve_instance_dir(cls, instance_name: str) -> str: + """ + Resolve the `bots/instances` directory that belongs to `instance_name`. + + Bot containers are named after their instance verbatim, so this directory is also what + identifies a container as one this API created. Raises ValueError if the name escapes + `bots/instances`. + """ + instances_base = os.path.join("bots", "instances") + instance_dir = os.path.join(instances_base, instance_name) + cls._ensure_contained(instance_dir, instances_base, "instance_name") + return instance_dir + def create_hummingbot_instance(self, config: V2ControllerDeployment): bots_path = os.environ.get('BOTS_PATH', self.SOURCE_PATH) # Default to 'SOURCE_PATH' if BOTS_PATH is not set instance_name = config.instance_name @@ -226,10 +281,29 @@ def create_hummingbot_instance(self, config: V2ControllerDeployment): os.makedirs(destination_controllers_config_dir, exist_ok=True) for controller_file in controllers_list: - source_controller_file = os.path.join(controllers_config_dir, controller_file) - destination_controller_file = os.path.join( - destination_controllers_config_dir, controller_file - ) + # SEC-058: the controllers list is read back from an attacker-controllable + # YAML file, so it never went through the request-body validators. Validate + # each entry as a single safe path component and, as defense in depth, + # verify both resolved paths stay inside their base directories. + try: + if not isinstance(controller_file, str): + raise ValueError( + f"Invalid controllers_config entry: {controller_file!r} is not a string." + ) + validate_safe_config_name(controller_file, "controllers_config") + source_controller_file = self._ensure_contained( + os.path.join(controllers_config_dir, controller_file), + controllers_config_dir, + "controllers_config", + ) + destination_controller_file = self._ensure_contained( + os.path.join(destination_controllers_config_dir, controller_file), + destination_controllers_config_dir, + "controllers_config", + ) + except ValueError as e: + logger.warning(f"Skipping unsafe controller config entry {controller_file!r}: {e}") + continue if os.path.exists(source_controller_file): shutil.copy2(source_controller_file, destination_controller_file) diff --git a/services/executor_service.py b/services/executor_service.py index 955d4caf..4352e09e 100644 --- a/services/executor_service.py +++ b/services/executor_service.py @@ -7,7 +7,7 @@ import json import logging import time -from datetime import datetime, timezone +from datetime import datetime, timedelta, timezone from decimal import Decimal from enum import Enum from typing import Any, Dict, List, Optional, Type @@ -32,8 +32,15 @@ from hummingbot.strategy_v2.executors.xemm_executor.data_types import XEMMExecutorConfig from hummingbot.strategy_v2.executors.xemm_executor.xemm_executor import XEMMExecutor from hummingbot.strategy_v2.models.executors import CloseType, TrackedOrder - -from database import AsyncDatabaseManager, ExecutorRepository, GatewayCLMMRepository, GatewaySwapRepository +from sqlalchemy.exc import IntegrityError + +from database import ( + AsyncDatabaseManager, + ExecutorPerformanceRepository, + ExecutorRepository, + GatewayCLMMRepository, + GatewaySwapRepository, +) from models.executors import PositionHold from services.gateway_client import get_native_gas_token from services.trading_service import AccountTradingInterface, TradingService @@ -116,6 +123,11 @@ class ExecutorService: # poller creates it on discovery, which is far slower than the control loop's tick. LP_RENT_RETRY_SECONDS = 30.0 + # How often the retention sweep runs, once it is enabled at all. Deleting rows older + # than N days does not get more correct for being done every minute, and the sweep is + # a DELETE over two tables sharing the control loop's tick. + PERFORMANCE_PRUNE_INTERVAL_SECONDS = 3600.0 + # Mapping of executor type strings to (executor_class, config_class) EXECUTOR_REGISTRY: Dict[str, tuple[Type[ExecutorBase], Type[ExecutorConfigBase]]] = { "position_executor": (PositionExecutor, PositionExecutorConfig), @@ -134,7 +146,9 @@ def __init__( db_manager: AsyncDatabaseManager, default_account: str = "master_account", update_interval: float = 1.0, - max_retries: int = 10 + max_retries: int = 10, + performance_snapshot_interval: float = 60.0, + performance_retention_days: int = 0 ): """ Initialize ExecutorService. @@ -145,12 +159,18 @@ def __init__( default_account: Default account to use update_interval: Executor update interval in seconds max_retries: Maximum retries for executor operations + performance_snapshot_interval: How often a live executor's performance is + written to executor_performance_snapshots, in seconds + performance_retention_days: Delete performance snapshots older than this many + days; 0 (the default) keeps everything """ self._trading_service = trading_service self.db_manager = db_manager self.default_account = default_account self.update_interval = update_interval self.max_retries = max_retries + self.performance_snapshot_interval = performance_snapshot_interval + self.performance_retention_days = performance_retention_days # Trading interfaces per account (lazy initialized via TradingService) self._trading_interfaces: Dict[str, AccountTradingInterface] = {} @@ -182,6 +202,11 @@ def __init__( # would be a query a second against a row that appears about once a minute. self._lp_rent_retry_after: Dict[str, float] = {} + # Performance snapshot cadence, on the control loop's own clock (monotonic, so a + # system clock step cannot stall the series or flood it). + self._last_snapshot_at: float = 0.0 + self._last_prune_at: float = 0.0 + # Control loop task self._control_loop_task: Optional[asyncio.Task] = None self._is_running = False @@ -339,6 +364,23 @@ async def _control_loop(self): for executor_id in completed_ids: await self._handle_executor_completion(executor_id) + # Performance snapshots, after completion handling so a just-closed + # executor is already out of _active_executors and gets only its terminal + # row, not a duplicate periodic one. This is not a separate task on + # purpose: the loop already ticks at 1 Hz and already awaits the database + # inside the tick (see _record_lp_position_rent), so the cadence is one + # guard rather than another thing to schedule and shut down. + now = time.monotonic() + if now - self._last_snapshot_at >= self.performance_snapshot_interval: + self._last_snapshot_at = now + await self._dump_executor_performance() + if ( + self.performance_retention_days > 0 + and now - self._last_prune_at >= self.PERFORMANCE_PRUNE_INTERVAL_SECONDS + ): + self._last_prune_at = now + await self._prune_performance_snapshots() + except Exception as e: logger.error(f"Error in executor control loop: {e}", exc_info=True) @@ -540,6 +582,179 @@ async def _record_executor_swap(self, executor_id: str, executor: ExecutorBase) except Exception as e: logger.error(f"Error recording executor swap {transaction_hash}: {e}", exc_info=True) + # ======================================== + # Performance Snapshots + # ======================================== + + def _build_snapshot_row( + self, + executor_id: str, + executor: ExecutorBase, + *, + is_terminal: bool, + metrics: Optional[Dict[str, Any]] = None, + status: Optional[str] = None, + close_type: Optional[str] = None, + ) -> Optional[Dict[str, Any]]: + """One executor_performance_snapshots row, from the two sources _format_executor_info reads. + + `metrics`, `status` and `close_type` let the terminal row reuse the figures + _persist_executor_completed already computed rather than reading executor_info a + second time -- the read that can raise and get silently substituted with zeros. + + Returns None when the metrics cannot be read at all: a row of fabricated zeros in + the middle of a series is worse than a gap, because a reader cannot tell it from + an executor that genuinely made nothing. + """ + metadata = self._executor_metadata.get(executor_id, {}) + + if metrics is None: + try: + info = executor.executor_info + metrics = { + "net_pnl_quote": info.net_pnl_quote, + "net_pnl_pct": info.net_pnl_pct, + "cum_fees_quote": info.cum_fees_quote, + "filled_amount_quote": info.filled_amount_quote, + } + except Exception as e: + logger.debug(f"Could not read executor_info for {executor_id} while snapshotting: {e}") + return None + + if status is None: + try: + status = executor.status.name + except Exception as e: + logger.debug(f"Could not read status for {executor_id} while snapshotting: {e}") + return None + + return { + "executor_id": executor_id, + "executor_type": metadata.get("executor_type") or "unknown", + "account_name": metadata.get("account_name") or self.default_account, + "connector_name": metadata.get("connector_name") or "", + "trading_pair": metadata.get("trading_pair") or "", + "controller_id": metadata.get("controller_id", "main"), + "status": status, + "close_type": close_type, + "is_terminal": is_terminal, + **metrics, + } + + async def _dump_executor_performance(self): + """Write one snapshot row per live executor. + + Only live executors are sampled because _active_executors is, by construction, + the only place a running executor's performance exists -- its database row still + carries the creation-time zeros until it completes. A closed executor is not + missing from here: it got its terminal row at completion. + + Never lets a database failure out: this shares the control loop's tick, and a + dropped snapshot must not stop executors from being updated. + """ + if not self.db_manager or not self._active_executors: + return + + snapshot_timestamp = datetime.now(timezone.utc) + rows = [] + for executor_id, executor in list(self._active_executors.items()): + row = self._build_snapshot_row(executor_id, executor, is_terminal=False) + if row is not None: + row["snapshot_timestamp"] = snapshot_timestamp + rows.append(row) + + if not rows: + return + + try: + async with self.db_manager.get_session_context() as session: + await ExecutorPerformanceRepository(session).save_snapshots(rows) + logger.debug(f"Dumped {len(rows)} executor performance snapshots") + except Exception as e: + logger.error(f"Error saving executor performance snapshots: {e}", exc_info=True) + + async def _prune_performance_snapshots(self): + """Delete snapshots older than the configured retention, from both snapshot tables. + + Gated on performance_retention_days > 0 by the caller: the default keeps + everything, so an upgrade never starts deleting an operator's history. + """ + if not self.db_manager or self.performance_retention_days <= 0: + return + + cutoff = datetime.now(timezone.utc) - timedelta(days=self.performance_retention_days) + try: + async with self.db_manager.get_session_context() as session: + executor_rows, controller_rows = await ExecutorPerformanceRepository( + session + ).prune_older_than(cutoff) + if executor_rows or controller_rows: + logger.info( + f"Pruned performance snapshots older than {cutoff.isoformat()}: " + f"{executor_rows} executor, {controller_rows} controller" + ) + except Exception as e: + logger.error(f"Error pruning performance snapshots: {e}", exc_info=True) + + async def get_executor_performance_history( + self, + executor_id: Optional[str] = None, + executor_type: Optional[str] = None, + controller_id: Optional[str] = None, + account_name: Optional[str] = None, + connector_name: Optional[str] = None, + trading_pair: Optional[str] = None, + limit: Optional[int] = None, + cursor: Optional[str] = None, + start_time: Optional[datetime] = None, + end_time: Optional[datetime] = None, + interval: str = "5m" + ): + """Read an executor snapshot series. Mirrors BotsOrchestrator.get_controller_performance_history.""" + async with self.db_manager.get_session_context() as session: + repo = ExecutorPerformanceRepository( + session, grain_minutes=self.performance_snapshot_interval / 60.0 + ) + return await repo.get_performance_history( + executor_id=executor_id, + executor_type=executor_type, + controller_id=controller_id, + account_name=account_name, + connector_name=connector_name, + trading_pair=trading_pair, + limit=limit, + cursor=cursor, + start_time=start_time, + end_time=end_time, + interval=interval, + ) + + async def get_latest_executor_performance( + self, + executor_id: Optional[str] = None, + executor_type: Optional[str] = None, + controller_id: Optional[str] = None, + account_name: Optional[str] = None, + connector_name: Optional[str] = None, + trading_pair: Optional[str] = None, + limit: Optional[int] = None, + ): + """The last snapshot of each matching executor. Mirrors + BotsOrchestrator.get_latest_controller_performance.""" + async with self.db_manager.get_session_context() as session: + repo = ExecutorPerformanceRepository( + session, grain_minutes=self.performance_snapshot_interval / 60.0 + ) + return await repo.get_latest( + executor_id=executor_id, + executor_type=executor_type, + controller_id=controller_id, + account_name=account_name, + connector_name=connector_name, + trading_pair=trading_pair, + limit=limit, + ) + def _get_trading_interface(self, account_name: str) -> AccountTradingInterface: """Get or create an AccountTradingInterface for the account.""" if account_name not in self._trading_interfaces: @@ -1247,8 +1462,6 @@ async def get_performance_report( Returns: Dictionary with performance metrics ready for PerformanceReportResponse. """ - import math - report: Dict[str, Any] = { "controller_id": controller_id, "total_executors": 0, @@ -1266,31 +1479,33 @@ async def get_performance_report( } if self.db_manager: - try: - async with self.db_manager.get_session_context() as session: - repo = ExecutorRepository(session) - db_data = await repo.get_performance_report(controller_id=controller_id) - - report["total_executors"] = db_data["total_executors"] - report["by_status"] = db_data["status_counts"] - report["pnl_total_quote"] = db_data["pnl_total_quote"] - report["pnl_pct_avg"] = db_data["pnl_pct_avg"] - report["fees_total_quote"] = db_data["fees_total_quote"] - report["volume_total_quote"] = db_data["volume_total_quote"] - report["win_rate"] = db_data["win_rate"] - report["by_type"] = db_data["by_type"] - - # Sharpe ratio: mean(pnl) / std(pnl), requires >= 2 values - pnl_values = db_data.get("pnl_values", []) - if len(pnl_values) >= 2: - mean_pnl = sum(pnl_values) / len(pnl_values) - variance = sum((v - mean_pnl) ** 2 for v in pnl_values) / (len(pnl_values) - 1) - std_pnl = math.sqrt(variance) - if std_pnl > 0: - report["sharpe_ratio"] = round(mean_pnl / std_pnl, 4) - - except Exception as e: - logger.error(f"Error generating performance report: {e}", exc_info=True) + # No try/except here on purpose: a database outage used to be swallowed and + # reported as the zeroed report below, making it indistinguishable from + # "no executors yet" to every consumer -- the route answered 200 with zeroes + # and the /ws/executors performance channel pushed a confident row of them. + # The failure now propagates: the route turns it into a 500 and the push loop + # sends an error frame. The zeroed report is left to mean an empty dataset. + async with self.db_manager.get_session_context() as session: + repo = ExecutorRepository(session) + db_data = await repo.get_performance_report(controller_id=controller_id) + + report["total_executors"] = db_data["total_executors"] + report["by_status"] = db_data["status_counts"] + report["pnl_total_quote"] = db_data["pnl_total_quote"] + report["pnl_pct_avg"] = db_data["pnl_pct_avg"] + report["fees_total_quote"] = db_data["fees_total_quote"] + report["volume_total_quote"] = db_data["volume_total_quote"] + report["win_rate"] = db_data["win_rate"] + report["by_type"] = db_data["by_type"] + + # Sharpe ratio: mean(pnl) / std(pnl). The repository returns the sample + # standard deviation as an aggregate -- None below 2 completed executors -- + # so this never depends on the number of rows in the table. + pnl_std = db_data["pnl_std"] + completed_count = db_data["completed_count"] + if pnl_std: # None below 2 executors, 0.0 when every PnL is identical + mean_pnl = db_data["pnl_total_quote"] / completed_count + report["sharpe_ratio"] = round(mean_pnl / pnl_std, 4) # --- Unrealized PnL from active executors --- unrealized_pnl = 0.0 @@ -1379,6 +1594,14 @@ async def _persist_executor_created(self, executor_id: str, executor: ExecutorBa logger.debug(f"Persisted executor {executor_id} creation to database") + except IntegrityError: + # The executor closed before this INSERT landed and its completion already + # wrote the row (see upsert_executor_completion). That row carries the final + # state; re-inserting a RUNNING one is exactly what must not happen. + logger.debug( + f"Executor {executor_id} row already written by its completion; " + f"skipping the creation insert" + ) except Exception as e: logger.error(f"Error persisting executor creation: {e}") @@ -1399,12 +1622,18 @@ async def _persist_executor_completed(self, executor_id: str, executor: Executor net_pnl_pct = executor_info.net_pnl_pct cum_fees_quote = executor_info.cum_fees_quote filled_amount_quote = executor_info.filled_amount_quote + metrics_are_measured = True except Exception as e: logger.debug(f"Error accessing executor_info for persistence: {e}") net_pnl_quote = Decimal("0") net_pnl_pct = Decimal("0") cum_fees_quote = Decimal("0") filled_amount_quote = Decimal("0") + # The record still takes these zeros, as it always has. The snapshot + # series does not: a terminal row is what a reader takes for the + # executor's final value, and a fabricated zero there is indistinguishable + # from an executor that genuinely made nothing. + metrics_are_measured = False # Get custom_info directly from executor to avoid Pydantic serialization issues # with TrackedOrder and other complex types @@ -1465,8 +1694,19 @@ async def _persist_executor_completed(self, executor_id: str, executor: Executor async with self.db_manager.get_session_context() as session: repo = ExecutorRepository(session) - await repo.update_executor( + # Upsert, not update: an executor that closes in milliseconds can reach + # here before _persist_executor_created's INSERT has landed (or after it + # failed outright), and a plain select-then-update would silently drop + # this final state, leaving a phantom RUNNING executor forever. + record, repaired = await repo.upsert_executor_completion( executor_id=executor_id, + executor_type=executor_type, + account_name=metadata.get("account_name"), + connector_name=metadata.get("connector_name"), + trading_pair=metadata.get("trading_pair"), + controller_id=metadata.get("controller_id", "main"), + config=json.dumps(metadata.get("config", {}), default=_json_default), + created_at=metadata.get("created_at"), status=status_name, close_type=close_type, net_pnl_quote=net_pnl_quote, @@ -1477,7 +1717,54 @@ async def _persist_executor_completed(self, executor_id: str, executor: Executor error_log=error_log_json ) - logger.debug(f"Persisted executor {executor_id} completion to database") + # The terminal row, in the same transaction as the record update so the + # two never disagree. It is what makes a closed executor's series + # answerable from executor_performance_snapshots alone -- no join, and no + # "and then append the final value" rule for every future reader. + # + # Behind a SAVEPOINT, because the direction of the coupling matters: the + # record update is the accounting and the snapshot is a point on a chart, + # so a failing snapshot must roll back only itself. Without it a bad + # INSERT here would abort the transaction and lose the completion, which + # is the failure 60041a8 went to some trouble to make impossible. + if record is not None and metrics_are_measured: + terminal_row = self._build_snapshot_row( + executor_id, + executor, + is_terminal=True, + metrics={ + "net_pnl_quote": net_pnl_quote, + "net_pnl_pct": net_pnl_pct, + "cum_fees_quote": cum_fees_quote, + "filled_amount_quote": filled_amount_quote, + }, + status=status_name, + close_type=close_type, + ) + if terminal_row is not None: + try: + async with session.begin_nested(): + await ExecutorPerformanceRepository(session).save_snapshots( + [terminal_row] + ) + except Exception as e: + logger.error( + f"Could not write the terminal performance snapshot for " + f"{executor_id}; its series ends at the last periodic row: {e}" + ) + + if record is None: + logger.error( + f"Could not persist completion for executor {executor_id}: no row to " + f"update and the repair insert did not take" + ) + elif repaired: + logger.warning( + f"Executor {executor_id} completed before its creation row existed; " + f"inserted the record from its final state" + ) + else: + logger.debug(f"Persisted executor {executor_id} completion to database") except Exception as e: logger.error(f"Error persisting executor completion: {e}") diff --git a/services/executor_ws_manager.py b/services/executor_ws_manager.py index 29191925..9be30af1 100644 --- a/services/executor_ws_manager.py +++ b/services/executor_ws_manager.py @@ -10,10 +10,13 @@ import logging import time from dataclasses import dataclass, field -from typing import Any, Dict, Optional +from functools import partial +from typing import Any, Awaitable, Callable, Dict, Optional from fastapi import WebSocket +from fastapi.websockets import WebSocketDisconnect +from config import settings from services.bots_orchestrator import BotsOrchestrator from services.executor_service import ExecutorService from services.market_data_service import MarketDataService @@ -21,11 +24,6 @@ logger = logging.getLogger(__name__) -# Update interval bounds (seconds) -MIN_UPDATE_INTERVAL = 0.5 -MAX_UPDATE_INTERVAL = 60.0 -DEFAULT_UPDATE_INTERVAL = 2.0 - SUBSCRIPTION_TYPES = { "executors", "executor_detail", @@ -66,6 +64,23 @@ class ExecutorSubscription: last_sent_hash: Optional[str] = None # For logs: track count to send only new entries last_log_count: int = 0 + # Whether the client has already been told this channel is failing, so a + # persistent fault sends one error frame instead of one per interval. + error_notified: bool = False + + +# A fetcher turns a subscription into the payload to hash and push; an extra +# builder derives additional top-level frame keys from that payload. +FetchFn = Callable[["ExecutorSubscription"], Awaitable[Any]] +ExtraFn = Callable[[Any], Dict[str, Any]] + + +@dataclass(frozen=True) +class PushSpec: + """How one hash-and-push subscription type is fetched and framed.""" + fetch: FetchFn + msg_type: str + extra: Optional[ExtraFn] = None def _compute_hash(data: Any) -> str: @@ -75,10 +90,25 @@ def _compute_hash(data: Any) -> str: def _clamp_interval(interval: Optional[float]) -> float: - """Clamp update interval to allowed range.""" + """Clamp update interval to the configured executor WebSocket range. + + Raises ValueError when the client-supplied value is not a number. The + comparisons below would otherwise raise TypeError out of handle_subscribe, + surfacing as an unhandled exception instead of the error frame every other + malformed-input path returns (CORR-113). + """ + md = settings.market_data if interval is None: - return DEFAULT_UPDATE_INTERVAL - return max(MIN_UPDATE_INTERVAL, min(MAX_UPDATE_INTERVAL, interval)) + return md.ws_executor_default_update_interval + # bool is an int subclass, so True would otherwise clamp to the floor. + if isinstance(interval, bool) or not isinstance(interval, (int, float)): + raise ValueError( + f"update_interval must be a number, got {type(interval).__name__}" + ) + return max( + md.ws_executor_min_update_interval, + min(md.ws_executor_max_update_interval, interval), + ) class ExecutorWebSocketManager: @@ -115,7 +145,11 @@ async def handle_subscribe( ) return - interval = _clamp_interval(msg.get("update_interval")) + try: + interval = _clamp_interval(msg.get("update_interval")) + except ValueError as e: + await self._send_error(websocket, str(e)) + return # Build subscription sub = ExecutorSubscription( @@ -240,198 +274,101 @@ async def shutdown(self) -> None: # Push loop dispatch # ------------------------------------------------------------------ - def _get_push_fn(self, sub_type: str): + def _push_specs(self) -> Dict[str, PushSpec]: + """Map every hash-and-push subscription type to how it is fetched and framed.""" return { - "executors": self._executors_push_loop, - "executor_detail": self._executor_detail_push_loop, - "executor_summary": self._summary_push_loop, - "performance": self._performance_push_loop, - "positions": self._positions_push_loop, - "executor_logs": self._logs_push_loop, - "bot_status": self._bot_status_push_loop, - "all_bots_status": self._all_bots_status_push_loop, - }[sub_type] + "executors": PushSpec( + fetch=self._fetch_executors, + msg_type="executors", + extra=lambda data: {"total_count": len(data)}, + ), + "executor_detail": PushSpec( + fetch=self._fetch_executor_detail, + msg_type="executor_detail", + ), + "executor_summary": PushSpec( + fetch=self._fetch_summary, + msg_type="executor_summary", + ), + "performance": PushSpec( + fetch=self._fetch_performance, + msg_type="performance", + ), + "positions": PushSpec( + fetch=self._fetch_positions, + msg_type="positions", + ), + "bot_status": PushSpec( + fetch=self._fetch_bot_status, + msg_type="bot_status", + ), + "all_bots_status": PushSpec( + fetch=self._fetch_all_bots_status, + msg_type="all_bots_status", + extra=lambda data: {"bot_count": len(data)}, + ), + } + + def _get_push_fn(self, sub_type: str): + """Resolve a subscription type to the coroutine function that drives its loop.""" + if sub_type == "executor_logs": + # Logs key on last_log_count, not on a payload hash — its own loop. + return self._logs_push_loop + spec = self._push_specs()[sub_type] + return partial( + self._push_loop, + fetch=spec.fetch, + msg_type=spec.msg_type, + extra=spec.extra, + ) # ------------------------------------------------------------------ # Push loops # ------------------------------------------------------------------ - async def _executors_push_loop( - self, conn_id: str, websocket: WebSocket, sub: ExecutorSubscription - ) -> None: - """Poll get_executors() with filters and push on change.""" - try: - while True: - try: - filters = sub.filters - executors = await self._executor_service.get_executors( - account_name=filters.get("account_name"), - connector_name=filters.get("connector_name"), - trading_pair=filters.get("trading_pair"), - executor_type=filters.get("executor_type"), - status=filters.get("status"), - controller_id=filters.get("controller_id"), - ) - h = _compute_hash(executors) - if h != sub.last_sent_hash: - sub.last_sent_hash = h - await websocket.send_json({ - "type": "executors", - "subscription_id": sub.sub_id, - "data": executors, - "total_count": len(executors), - "timestamp": time.time(), - }) - except Exception as e: - logger.error(f"[WS-Exec] executors push error: {e}", exc_info=True) - await asyncio.sleep(sub.update_interval) - except asyncio.CancelledError: - pass - - async def _executor_detail_push_loop( - self, conn_id: str, websocket: WebSocket, sub: ExecutorSubscription - ) -> None: - """Poll get_executor() for a single executor and push on change.""" - try: - while True: - try: - data = await self._executor_service.get_executor(sub.executor_id) - h = _compute_hash(data) - if h != sub.last_sent_hash: - sub.last_sent_hash = h - await websocket.send_json({ - "type": "executor_detail", - "subscription_id": sub.sub_id, - "data": data, - "timestamp": time.time(), - }) - except Exception as e: - logger.error(f"[WS-Exec] executor_detail push error: {e}", exc_info=True) - await asyncio.sleep(sub.update_interval) - except asyncio.CancelledError: - pass - - async def _summary_push_loop( - self, conn_id: str, websocket: WebSocket, sub: ExecutorSubscription - ) -> None: - """Poll get_summary() and push on change.""" - try: - while True: - try: - data = self._executor_service.get_summary() - h = _compute_hash(data) - if h != sub.last_sent_hash: - sub.last_sent_hash = h - await websocket.send_json({ - "type": "executor_summary", - "subscription_id": sub.sub_id, - "data": data, - "timestamp": time.time(), - }) - except Exception as e: - logger.error(f"[WS-Exec] summary push error: {e}", exc_info=True) - await asyncio.sleep(sub.update_interval) - except asyncio.CancelledError: - pass - - async def _performance_push_loop( - self, conn_id: str, websocket: WebSocket, sub: ExecutorSubscription + async def _push_loop( + self, + conn_id: str, + websocket: WebSocket, + sub: ExecutorSubscription, + fetch: FetchFn, + msg_type: str, + extra: Optional[ExtraFn] = None, ) -> None: - """Poll get_performance_report() and push on change.""" + """Poll `fetch` on the subscription interval and push only when the data changes.""" try: while True: try: - data = await self._executor_service.get_performance_report( - controller_id=sub.controller_id, - market_data_service=self._market_data_service, - ) + data = await fetch(sub) + sub.error_notified = False h = _compute_hash(data) if h != sub.last_sent_hash: sub.last_sent_hash = h - await websocket.send_json({ - "type": "performance", + message = { + "type": msg_type, "subscription_id": sub.sub_id, "data": data, - "timestamp": time.time(), - }) - except Exception as e: - logger.error(f"[WS-Exec] performance push error: {e}", exc_info=True) - await asyncio.sleep(sub.update_interval) - except asyncio.CancelledError: - pass - - async def _positions_push_loop( - self, conn_id: str, websocket: WebSocket, sub: ExecutorSubscription - ) -> None: - """Poll get_positions_held() with unrealized PnL and push on change.""" - try: - while True: - try: - positions = self._executor_service.get_positions_held( - controller_id=sub.controller_id, - ) - # Build response dicts with unrealized PnL - position_dicts = [] - total_realized = 0.0 - total_unrealized = None - - for p in positions: - unrealized_pnl = None - # See routers/executors.py: a hyphenated base symbol failed the - # old length check and dropped the PnL without saying so. - try: - base, quote = split_trading_pair(p.trading_pair) - except InvalidTradingPair: - base = quote = None - if base and quote: - rate = self._market_data_service.get_rate(base, quote) - if rate is not None: - unrealized_pnl = float(p.get_unrealized_pnl(rate)) - if total_unrealized is None: - total_unrealized = 0.0 - total_unrealized += unrealized_pnl - - total_realized += float(p.realized_pnl_quote) - position_dicts.append({ - "trading_pair": p.trading_pair, - "connector_name": p.connector_name, - "account_name": p.account_name, - "controller_id": p.controller_id, - "buy_amount_base": float(p.buy_amount_base), - "buy_amount_quote": float(p.buy_amount_quote), - "sell_amount_base": float(p.sell_amount_base), - "sell_amount_quote": float(p.sell_amount_quote), - "net_amount_base": float(p.net_amount_base), - "buy_breakeven_price": float(p.buy_breakeven_price) if p.buy_breakeven_price else None, - "sell_breakeven_price": float(p.sell_breakeven_price) if p.sell_breakeven_price else None, - "matched_amount_base": float(p.matched_amount_base), - "unmatched_amount_base": float(p.unmatched_amount_base), - "position_side": p.position_side, - "realized_pnl_quote": float(p.realized_pnl_quote), - "unrealized_pnl_quote": unrealized_pnl, - "executor_count": len(p.executor_ids), - "executor_ids": p.executor_ids, - "last_updated": p.last_updated.isoformat() if p.last_updated else None, - }) - - payload = { - "total_positions": len(positions), - "total_realized_pnl": total_realized, - "total_unrealized_pnl": total_unrealized, - "positions": position_dicts, - } - - h = _compute_hash(payload) - if h != sub.last_sent_hash: - sub.last_sent_hash = h - await websocket.send_json({ - "type": "positions", - "subscription_id": sub.sub_id, - "data": payload, - "timestamp": time.time(), - }) + } + if extra is not None: + message.update(extra(data)) + message["timestamp"] = time.time() + if not await self._send_or_stop(conn_id, websocket, sub, msg_type, message): + break except Exception as e: - logger.error(f"[WS-Exec] positions push error: {e}", exc_info=True) + logger.error(f"[WS-Exec] {msg_type} push error: {e}", exc_info=True) + # Tell the subscriber the channel is failing. Staying silent leaves + # the last good frame on screen with nothing marking it stale, and + # a fetch that degrades to zeroes instead of raising would render as + # real data (see get_performance_report, CORR-111). + if not sub.error_notified: + sub.error_notified = True + # The recovery frame must be sent even if the payload is + # byte-identical to the one that preceded the outage. + sub.last_sent_hash = None + if not await self._send_or_stop( + conn_id, websocket, sub, msg_type, self._error_frame(sub, msg_type, e) + ): + break await asyncio.sleep(sub.update_interval) except asyncio.CancelledError: pass @@ -452,86 +389,177 @@ async def _logs_push_loop( if current_count > sub.last_log_count: new_logs = all_logs[sub.last_log_count:] sub.last_log_count = current_count - await websocket.send_json({ + message = { "type": "executor_logs", "subscription_id": sub.sub_id, "data": new_logs, "total_count": current_count, "timestamp": time.time(), - }) + } + if not await self._send_or_stop( + conn_id, websocket, sub, "executor_logs", message + ): + break except Exception as e: - logger.error(f"[WS-Exec] logs push error: {e}", exc_info=True) + logger.error(f"[WS-Exec] executor_logs push error: {e}", exc_info=True) await asyncio.sleep(sub.update_interval) except asyncio.CancelledError: pass - async def _bot_status_push_loop( - self, conn_id: str, websocket: WebSocket, sub: ExecutorSubscription - ) -> None: - """Poll get_bot_status() for a single bot and push on change (no logs).""" - try: - while True: - try: - raw_status = self._bots_orchestrator.get_bot_status(sub.bot_name) - # Strip logs — only send status, performance, custom_info - payload = { - "bot_name": sub.bot_name, - "status": raw_status.get("status"), - "performance": raw_status.get("performance", {}), - "recently_active": raw_status.get("recently_active", False), - } - h = _compute_hash(payload) - if h != sub.last_sent_hash: - sub.last_sent_hash = h - await websocket.send_json({ - "type": "bot_status", - "subscription_id": sub.sub_id, - "data": payload, - "timestamp": time.time(), - }) - except Exception as e: - logger.error(f"[WS-Exec] bot_status push error: {e}", exc_info=True) - await asyncio.sleep(sub.update_interval) - except asyncio.CancelledError: - pass + # ------------------------------------------------------------------ + # Fetchers + # + # One per hash-and-push subscription type. They normalise the mixed + # sync/async service calls and own any per-type payload shaping, so the + # generic loop above only ever sees "await fetch(sub) -> data". + # ------------------------------------------------------------------ - async def _all_bots_status_push_loop( - self, conn_id: str, websocket: WebSocket, sub: ExecutorSubscription - ) -> None: - """Poll get_all_bots_status() and push on change (no logs).""" - try: - while True: - try: - raw = self._bots_orchestrator.get_all_bots_status() - # Strip logs from each bot - payload = {} - for bot_name, bot_data in raw.items(): - payload[bot_name] = { - "status": bot_data.get("status"), - "source": bot_data.get("source"), - "performance": bot_data.get("performance", {}), - "recently_active": bot_data.get("recently_active", False), - } - h = _compute_hash(payload) - if h != sub.last_sent_hash: - sub.last_sent_hash = h - await websocket.send_json({ - "type": "all_bots_status", - "subscription_id": sub.sub_id, - "data": payload, - "bot_count": len(payload), - "timestamp": time.time(), - }) - except Exception as e: - logger.error(f"[WS-Exec] all_bots_status push error: {e}", exc_info=True) - await asyncio.sleep(sub.update_interval) - except asyncio.CancelledError: - pass + async def _fetch_executors(self, sub: ExecutorSubscription) -> Any: + filters = sub.filters + return await self._executor_service.get_executors( + account_name=filters.get("account_name"), + connector_name=filters.get("connector_name"), + trading_pair=filters.get("trading_pair"), + executor_type=filters.get("executor_type"), + status=filters.get("status"), + controller_id=filters.get("controller_id"), + ) + + async def _fetch_executor_detail(self, sub: ExecutorSubscription) -> Any: + return await self._executor_service.get_executor(sub.executor_id) + + async def _fetch_summary(self, sub: ExecutorSubscription) -> Any: + return self._executor_service.get_summary() + + async def _fetch_performance(self, sub: ExecutorSubscription) -> Any: + return await self._executor_service.get_performance_report( + controller_id=sub.controller_id, + market_data_service=self._market_data_service, + ) + + async def _fetch_positions(self, sub: ExecutorSubscription) -> Dict[str, Any]: + """Positions held, enriched with unrealized PnL from the current rates.""" + positions = self._executor_service.get_positions_held( + controller_id=sub.controller_id, + ) + position_dicts = [] + total_realized = 0.0 + total_unrealized = None + + for p in positions: + unrealized_pnl = None + # See routers/executors.py: a hyphenated base symbol failed the + # old length check and dropped the PnL without saying so. + try: + base, quote = split_trading_pair(p.trading_pair) + except InvalidTradingPair: + base = quote = None + if base and quote: + rate = self._market_data_service.get_rate(base, quote) + if rate is not None: + unrealized_pnl = float(p.get_unrealized_pnl(rate)) + if total_unrealized is None: + total_unrealized = 0.0 + total_unrealized += unrealized_pnl + + total_realized += float(p.realized_pnl_quote) + position_dicts.append({ + "trading_pair": p.trading_pair, + "connector_name": p.connector_name, + "account_name": p.account_name, + "controller_id": p.controller_id, + "buy_amount_base": float(p.buy_amount_base), + "buy_amount_quote": float(p.buy_amount_quote), + "sell_amount_base": float(p.sell_amount_base), + "sell_amount_quote": float(p.sell_amount_quote), + "net_amount_base": float(p.net_amount_base), + "buy_breakeven_price": float(p.buy_breakeven_price) if p.buy_breakeven_price else None, + "sell_breakeven_price": float(p.sell_breakeven_price) if p.sell_breakeven_price else None, + "matched_amount_base": float(p.matched_amount_base), + "unmatched_amount_base": float(p.unmatched_amount_base), + "position_side": p.position_side, + "realized_pnl_quote": float(p.realized_pnl_quote), + "unrealized_pnl_quote": unrealized_pnl, + "executor_count": len(p.executor_ids), + "executor_ids": p.executor_ids, + "last_updated": p.last_updated.isoformat() if p.last_updated else None, + }) + + return { + "total_positions": len(positions), + "total_realized_pnl": total_realized, + "total_unrealized_pnl": total_unrealized, + "positions": position_dicts, + } + + async def _fetch_bot_status(self, sub: ExecutorSubscription) -> Dict[str, Any]: + """Single bot status, with the logs stripped out.""" + raw_status = self._bots_orchestrator.get_bot_status(sub.bot_name) + return { + "bot_name": sub.bot_name, + "status": raw_status.get("status"), + "performance": raw_status.get("performance", {}), + "recently_active": raw_status.get("recently_active", False), + } + + async def _fetch_all_bots_status(self, sub: ExecutorSubscription) -> Dict[str, Any]: + """Every bot's status, with the logs stripped out of each.""" + raw = self._bots_orchestrator.get_all_bots_status() + return { + bot_name: { + "status": bot_data.get("status"), + "source": bot_data.get("source"), + "performance": bot_data.get("performance", {}), + "recently_active": bot_data.get("recently_active", False), + } + for bot_name, bot_data in raw.items() + } # ------------------------------------------------------------------ # Helpers # ------------------------------------------------------------------ + @staticmethod + async def _send_or_stop( + conn_id: str, + websocket: WebSocket, + sub: ExecutorSubscription, + msg_type: str, + message: Dict[str, Any], + ) -> bool: + """Push one frame; return False when the client is gone and the loop must stop. + + Mirrors services/websocket_manager.py, which breaks out of its push loops + on a disconnect instead of logging an error every interval. Only the send + is guarded: a RuntimeError raised by a fetch is a service fault, not a + dropped client, and must stay a logged-and-retried error. + """ + try: + await websocket.send_json(message) + return True + except (WebSocketDisconnect, RuntimeError): + logger.info( + f"[WS-Exec] {conn_id} disconnected, stopping {msg_type} push [{sub.sub_id}]" + ) + return False + + @staticmethod + def _error_frame( + sub: ExecutorSubscription, msg_type: str, error: Exception + ) -> Dict[str, Any]: + """The frame a push loop sends when its fetch failed. + + Carries the subscription and channel so a client showing several + subscriptions can mark the right one as failing rather than stale. + """ + return { + "type": "error", + "subscription_id": sub.sub_id, + "channel": msg_type, + "message": f"{msg_type} update failed: {error}", + "timestamp": time.time(), + } + @staticmethod async def _send_error(websocket: WebSocket, message: str) -> None: await websocket.send_json({"type": "error", "message": message}) diff --git a/services/gateway_amm_service.py b/services/gateway_amm_service.py new file mode 100644 index 00000000..6322f19a --- /dev/null +++ b/services/gateway_amm_service.py @@ -0,0 +1,280 @@ +"""Persistence for AMM liquidity writes and the DAMM v2 positions that carry identity. + +Everything the ``/gateway/amm/*`` routes know about the database lives here: the shape +of a ``gateway_amm_events`` row, the shape of a ``gateway_amm_positions`` row and the +bookkeeping an add or a remove applies to it, and the pagination envelopes over both +searches. The handlers keep only what is genuinely theirs — reading Gateway's response +and deciding what to answer the caller. + +Reads propagate their errors; the after-the-fact recording of an operation that is +already on-chain is best-effort, per ``RepositoryService``. +""" +import logging +from decimal import Decimal +from typing import Any, Dict, Optional + +from database.repositories import GatewayAMMRepository +from services.gateway_client import get_native_gas_token +from services.repository_service import RepositoryService + +logger = logging.getLogger(__name__) + + +class GatewayAMMService(RepositoryService): + """AMM event and position persistence for the AMM routes.""" + + repository_class = GatewayAMMRepository + + # ------------------------------------------------------------------ + # Writes + # ------------------------------------------------------------------ + + async def record_position_add( + self, + *, + gateway_client, + position_address: str, + pool_address: str, + connector: str, + network: str, + wallet_address: str, + base_amount_added: Any, + quote_amount_added: Any, + position_rent: Any, + price: Optional[float], + base_token_address: str, + quote_token_address: str, + ) -> None: + """Create or top up the DAMM v2 position row for a confirmed add.""" + base_added = base_amount_added or 0 + quote_added = quote_amount_added or 0 + + async def _fn(repo): + existing = await repo.get_position_by_address(position_address) + if existing: + await repo.add_to_position_amounts( + position_address=position_address, + base_delta=Decimal(str(base_added)), + quote_delta=Decimal(str(quote_added)), + entry_price=Decimal(str(price)) if price else None, + ) + if existing.status == "CLOSED": + existing.status = "OPEN" + existing.closed_at = None + else: + chain, network_name = network.split("-", 1) + base_symbol = await gateway_client.resolve_token_symbol( + chain, network_name, base_token_address) + quote_symbol = await gateway_client.resolve_token_symbol( + chain, network_name, quote_token_address) + await repo.create_position({ + "position_address": position_address, + "pool_address": pool_address, + "connector": connector.split("/")[0], + "network": network, + "wallet_address": wallet_address, + "base_token": base_symbol, + "quote_token": quote_symbol, + "trading_pair": f"{base_symbol}-{quote_symbol}", + "initial_base_token_amount": base_added, + "initial_quote_token_amount": quote_added, + "base_token_amount": base_added, + "quote_token_amount": quote_added, + # Rent is locked, not spent: the chain returns it when the position + # account closes, and Gateway reports it separately for exactly that + # reason. Recorded here so the close can be checked against it — a + # refund smaller than what was locked means an account was left + # behind. Present only when this add opened the position; adding to + # one that already exists locks no further rent. + "position_rent": Decimal(str(position_rent)) if position_rent is not None else None, + "entry_price": price, + "current_price": price, + }) + logger.info(f"Booked AMM position {position_address}: +{base_added} base, " + f"+{quote_added} quote") + + await self._in_repo_best_effort( + _fn, error_message=f"Error booking AMM position {position_address}") + + async def record_position_remove( + self, + *, + position_address: str, + base_amount_removed: Any, + quote_amount_removed: Any, + percentage_to_remove: float, + position_rent_refunded: Any, + ) -> None: + """Unbook a confirmed removal, closing the row when the whole position went out.""" + async def _fn(repo): + position = await repo.subtract_from_position_amounts( + position_address=position_address, + base_delta=Decimal(str(base_amount_removed or 0)), + quote_delta=Decimal(str(quote_amount_removed or 0)), + ) + # A 100% remove is the close: Gateway closes the position account in + # the same transaction, which is what returns its rent. There is no + # separate close route, and positionRentRefunded arrives only on this + # path — a partial removal leaves the account open and refunds + # nothing, so its absence there is a fact rather than a gap. + if position and percentage_to_remove >= 100: + await repo.close_position( + position_address, + position_rent_refunded=(Decimal(str(position_rent_refunded)) + if position_rent_refunded is not None else None), + ) + + await self._in_repo_best_effort( + _fn, error_message=f"Error booking AMM removal for {position_address}") + + async def record_event( + self, + *, + transaction_hash: str, + event_type: str, + connector: str, + network: str, + wallet_address: str, + pool_address: str, + position_address: Optional[str], + base_token_amount: Any, + quote_token_amount: Any, + price: Optional[float], + gas_fee: Any, + tx_status: str, + ) -> None: + """Persist one AMM write. + + Best-effort: the liquidity has already moved by the time this is called, so a + database problem must not surface as a failed write to the caller. Amounts come + from Gateway's ``data``, present only once it confirmed the tx; a + submitted-not-confirmed write records the status with null amounts rather than + inventing figures. + """ + chain, _ = network.split("-", 1) if "-" in network else (network, "") + + async def _fn(repo): + await repo.create_event({ + "transaction_hash": transaction_hash, + "connector": connector, + "network": network, + "wallet_address": wallet_address, + "pool_address": pool_address, + "position_address": position_address, + "event_type": event_type, + "base_token_amount": base_token_amount, + "quote_token_amount": quote_token_amount, + "price": price, + "gas_fee": gas_fee, + "gas_token": get_native_gas_token(chain) if gas_fee is not None else None, + "status": tx_status, + }) + logger.info(f"Recorded AMM {event_type}: {transaction_hash} (status: {tx_status})") + + await self._in_repo_best_effort( + _fn, error_message=f"Error recording AMM {event_type} event") + + async def record_failed_event( + self, + *, + transaction_hash: Optional[str], + error: Exception, + event_type: str, + connector: str, + network: str, + wallet_address: str, + pool_address: str, + position_address: Optional[str] = None, + ) -> None: + """Record a write that reached the chain and reverted, before the error is re-raised. + + :meth:`record_event` above only runs when Gateway *returns*. A transaction that + landed and reverted does not return: Gateway raises, the client turns it into a + GatewayError, and control skips the whole recording block. That is why every row + in both event tables read CONFIRMED with no error_message — not because nothing + had ever failed, but because a failure could not be written. + + Only failures carrying a transaction id are recorded (the caller parses it out of + the error): a pre-flight simulation failure never got one and cost nothing, while + a landed revert has one and paid gas. Recording never masks the original failure. + """ + if not transaction_hash: + return + + async def _fn(repo): + await repo.create_event({ + "transaction_hash": transaction_hash, + "connector": connector, + "network": network, + "wallet_address": wallet_address, + "pool_address": pool_address, + "position_address": position_address, + "event_type": event_type, + "status": "FAILED", + "error_message": str(error), + }) + logger.error( + f"AMM {event_type} {transaction_hash} landed on-chain and FAILED on {connector}/" + f"{network}; recorded. {error}" + ) + + await self._in_repo_best_effort( + _fn, error_message=f"Error recording failed AMM {event_type}") + + # ------------------------------------------------------------------ + # Reads + # ------------------------------------------------------------------ + + async def search_events( + self, + *, + connector: Optional[str] = None, + network: Optional[str] = None, + wallet_address: Optional[str] = None, + pool_address: Optional[str] = None, + event_type: Optional[str] = None, + status: Optional[str] = None, + limit: int = 50, + offset: int = 0, + ) -> Dict[str, Any]: + """Recorded AMM liquidity writes, newest first, as the endpoint's envelope.""" + async def _fn(repo): + events = await repo.search_events( + connector=connector, network=network, wallet_address=wallet_address, + pool_address=pool_address, event_type=event_type, status=status, + limit=min(limit, 1000), offset=offset, + ) + return { + "data": [repo.event_to_dict(event) for event in events], + "total_count": len(events), + "limit": limit, + "offset": offset, + } + + return await self._in_repo(_fn) + + async def search_positions( + self, + *, + connector: Optional[str] = None, + network: Optional[str] = None, + wallet_address: Optional[str] = None, + pool_address: Optional[str] = None, + status: Optional[str] = None, + limit: int = 50, + offset: int = 0, + ) -> Dict[str, Any]: + """Tracked AMM positions (Meteora DAMM v2 NFTs), newest first, as the envelope.""" + async def _fn(repo): + positions = await repo.search_positions( + connector=connector, network=network, wallet_address=wallet_address, + pool_address=pool_address, status=status, limit=min(limit, 1000), offset=offset, + ) + return { + "data": [repo.position_to_dict(position) for position in positions], + "total_count": len(positions), + "limit": limit, + "offset": offset, + } + + return await self._in_repo(_fn) diff --git a/services/gateway_client.py b/services/gateway_client.py index 42eb4e48..678dcca3 100644 --- a/services/gateway_client.py +++ b/services/gateway_client.py @@ -1,5 +1,6 @@ import logging import ssl +import time from decimal import Decimal from typing import Any, Callable, Dict, List, Optional @@ -156,6 +157,13 @@ class GatewayClient: Provides essential functionality for wallet management and balance queries. """ + # How long a ``ping`` verdict stays reusable. The availability guard + # (``deps.require_gateway_online``) runs ahead of every guarded request, so without a cache + # each one pays a full round-trip to Gateway. The guard exists to catch "Gateway is down", + # not to timestamp the exact moment it went down, so a couple of seconds of staleness is the + # right trade — and any call that fails to connect clears the cache anyway (PERF-114). + PING_CACHE_TTL_SECONDS = 2.0 + def __init__( self, base_url: str = "http://localhost:15888", @@ -183,6 +191,8 @@ def __init__( self._connector_trading_types: Optional[Dict[str, List[str]]] = None # Per-(chain, network) token address -> symbol map, fetched on first use. self._token_symbols: Dict[tuple[str, str], Dict[str, str]] = {} + # Last ``ping`` verdict as (monotonic expiry, is_online), or None when unknown/invalidated. + self._ping_cache: Optional[tuple[float, bool]] = None @staticmethod def parse_network_id(network_id: str) -> tuple[str, str]: @@ -259,6 +269,7 @@ async def _request(self, method: str, path: str, params: Dict = None, json: Dict self._certs_unavailable_warned = True else: logger.debug(f"Gateway mTLS certs still unavailable, cannot reach {url}: {e}") + self._invalidate_ping_cache() return {"error": "Gateway client certificates not available; start the Gateway first", "status": 503} try: @@ -285,6 +296,9 @@ async def _request(self, method: str, path: str, params: Dict = None, json: Dict return await response.json() except aiohttp.ClientError as e: logger.debug(f"Gateway request error: {method} {url} - {e}") + # Gateway is unreachable: drop any cached "available" verdict so the next guarded + # request re-pings instead of waiting out the TTL on a stale answer. + self._invalidate_ping_cache() return None except Exception as e: logger.debug(f"Gateway request failed: {method} {url} - {e}") @@ -308,13 +322,24 @@ async def _get_error_body(self, response: aiohttp.ClientResponse) -> tuple: except Exception: return (f"HTTP {response.status}", None) + def _invalidate_ping_cache(self) -> None: + """Forget the cached availability verdict; the next ``ping`` hits Gateway again.""" + self._ping_cache = None + async def ping(self) -> bool: - """Check if Gateway is online""" + """Check if Gateway is online, reusing a verdict at most PING_CACHE_TTL_SECONDS old.""" + cached = self._ping_cache + if cached is not None and time.monotonic() < cached[0]: + return cached[1] try: response = await self._request("GET", "") - return response.get("status") == "ok" + online = response.get("status") == "ok" except Exception: - return False + online = False + # Assigned after the request so the invalidation _request performs on a connection + # failure cannot outrun the verdict it caused. + self._ping_cache = (time.monotonic() + self.PING_CACHE_TTL_SECONDS, online) + return online async def get_wallets(self) -> List[Dict]: """Get all connected wallets""" diff --git a/services/gateway_clmm_service.py b/services/gateway_clmm_service.py new file mode 100644 index 00000000..b33b5a99 --- /dev/null +++ b/services/gateway_clmm_service.py @@ -0,0 +1,1092 @@ +"""Persistence for CLMM positions and their lifecycle events. + +Everything the ``/gateway/clmm/*`` routes and the transaction poller know about the +database lives here: the shape of a position row, the shape of each event row (OPEN, +ADD_LIQUIDITY, REMOVE_LIQUIDITY, CLOSE, COLLECT_FEES), and the bookkeeping each one +applies to the position it belongs to. The handlers used to carry a copy of all three +per endpoint, which is how the poller's auto-discovery path ended up deriving a +different key set for the same table — see :func:`build_position_row`. + +Reads propagate their errors; the after-the-fact recording of an operation that is +already on-chain is best-effort, per ``RepositoryService``. +""" +import asyncio +import logging +from decimal import Decimal +from typing import Any, Dict, List, Optional, Set + +from database.repositories import GatewayCLMMRepository +from services.gateway_client import check_gateway_error +from services.repository_service import RepositoryService + +logger = logging.getLogger(__name__) + + +def build_position_row( + *, + position_address: str, + pool_address: str, + network: str, + connector: str, + wallet_address: str, + trading_pair: str, + base_token: str, + quote_token: str, + lower_price: Any, + upper_price: Any, + entry_price: Optional[Any], + current_price: Optional[Any], + initial_base_token_amount: Any, + initial_quote_token_amount: Any, + base_token_amount: Any, + quote_token_amount: Any, + lower_bin_id: Optional[int] = None, + upper_bin_id: Optional[int] = None, + position_rent: Optional[Any] = None, + in_range: str = "UNKNOWN", + base_fee_pending: Any = 0, + quote_fee_pending: Any = 0, +) -> Dict[str, Any]: + """The one description of what a new ``gateway_clmm_positions`` row contains. + + A position reaches the table by two routes — ``/gateway/clmm/open`` opens one, and + the poller's discovery sweep finds one that was opened elsewhere — and each used to + assemble the row itself. The key sets drifted apart: discovery wrote ``lower_bin_id``, + ``upper_bin_id``, ``base_fee_pending`` and ``quote_fee_pending``, the route wrote + ``position_rent``, and neither wrote the other's columns. The same logical position + therefore had two different row shapes depending on which path recorded it, and + anything reading the table back had to know which one to expect. + + So the key set is fixed here, and it is the union: every caller yields the same + columns, and a caller that cannot know a value passes nothing rather than omitting + the column. Absent is expressed as NULL (bin ids the route never learns, rent no + discovered position ever locked through us) and zero only where zero is the fact — + a position that has just come into the table has collected no fees. + + ``lower_price``/``upper_price``/the amounts accept anything ``float()`` takes, so + both a route's ``Decimal`` and Gateway's parsed JSON floats arrive the same way. + """ + # (upper - lower) / lower, computed on Decimals so a route's exact request values + # and the poller's floats round identically. + percentage = None + if lower_price and upper_price and Decimal(str(lower_price)) > 0: + percentage = float( + (Decimal(str(upper_price)) - Decimal(str(lower_price))) / Decimal(str(lower_price)) + ) + logger.info(f"Position price range percentage: {percentage:.4f} ({percentage * 100:.2f}%)") + + return { + "position_address": position_address, + "pool_address": pool_address, + "network": network, + "connector": connector, + "wallet_address": wallet_address, + "trading_pair": trading_pair, + "base_token": base_token, + "quote_token": quote_token, + "status": "OPEN", + "lower_price": float(lower_price), + "upper_price": float(upper_price), + # Bin-based CLMMs (Meteora) identify a range by bin; Gateway reports them on a + # position it lists, not on the response to opening one. + "lower_bin_id": lower_bin_id, + "upper_bin_id": upper_bin_id, + "entry_price": float(entry_price) if entry_price is not None else None, + "current_price": float(current_price) if current_price is not None else None, + "percentage": percentage, + "initial_base_token_amount": float(initial_base_token_amount), + "initial_quote_token_amount": float(initial_quote_token_amount), + # Rent is locked, not spent, and only the open route observes the figure. NULL + # rather than 0: a stored 0.0 claims a measurement that came back empty, which + # nothing downstream can tell from rent that was never read. + "position_rent": float(position_rent) if position_rent else None, + "base_token_amount": float(base_token_amount), + "quote_token_amount": float(quote_token_amount), + "in_range": in_range, + # Fees already accrued on-chain but not yet collected. Zero on a fresh open; + # a discovered position may have been earning for days before we saw it. + "base_fee_pending": float(base_fee_pending), + "quote_fee_pending": float(quote_fee_pending), + # Nothing has been collected through us yet on either path, by definition. + "base_fee_collected": 0.0, + "quote_fee_collected": 0.0, + } + + +async def refresh_position_data(position, gateway_client, clmm_repo: GatewayCLMMRepository): + """ + Refresh position data from Gateway and update database. + + This updates: + - in_range status + - liquidity amounts + - pending fees + - position status (if closed externally) + """ + try: + # Get wallet address for the position + wallet_address = position.wallet_address + + # Get all positions for this pool and find our specific position + try: + # check_gateway_error is critical here: a Gateway HTTP error must raise (and skip + # the refresh) rather than flow onward and mark the position CLOSED below. + positions_list = check_gateway_error(await gateway_client.clmm_positions_owned( + connector=position.connector, + chain_network=position.network, # position.network is already in 'chain-network' format + wallet_address=wallet_address + )) + + # Find our specific position in the list + result = None + if isinstance(positions_list, list): + for pos in positions_list: + if pos.get("address") == position.position_address: + result = pos + break + + # Absent from a single positions-owned read: could be closed externally, + # could be a lagging RPC node. Closing is owned by the poller's + # consecutive-miss gate (and the zero-liquidity check below) so one + # refresh can never close a live position. + if result is None: + logger.info(f"Position {position.position_address} absent from positions-owned; " + "skipping update (poller's miss-gate owns close detection)") + return + + except Exception as e: + # If we can't fetch positions, log error but don't mark as closed + logger.error(f"Error fetching position from Gateway: {e}") + return + + # Extract current state + current_price = Decimal(str(result.get("price", 0))) + lower_price = Decimal(str(result.get("lowerPrice", 0))) if result.get("lowerPrice") else Decimal("0") + upper_price = Decimal(str(result.get("upperPrice", 0))) if result.get("upperPrice") else Decimal("0") + + # Calculate in_range status + in_range = "UNKNOWN" + if current_price > 0 and lower_price > 0 and upper_price > 0: + if lower_price <= current_price <= upper_price: + in_range = "IN_RANGE" + else: + in_range = "OUT_OF_RANGE" + + # Extract token amounts + base_token_amount = Decimal(str(result.get("baseTokenAmount", 0))) + quote_token_amount = Decimal(str(result.get("quoteTokenAmount", 0))) + + # Check if position has been closed (zero liquidity) + if base_token_amount == 0 and quote_token_amount == 0: + logger.info(f"Position {position.position_address} has zero liquidity, marking as CLOSED") + await clmm_repo.close_position(position.position_address) + return + + # Update liquidity amounts, in_range status, and current price + await clmm_repo.update_position_liquidity( + position_address=position.position_address, + base_token_amount=base_token_amount, + quote_token_amount=quote_token_amount, + in_range=in_range, + current_price=current_price + ) + + # Always write pending fees — 0 is a real value (e.g. right after an + # external collect); the old non-zero guard left stale pendings forever. + base_fee_pending = Decimal(str(result.get("baseFeeAmount", 0))) + quote_fee_pending = Decimal(str(result.get("quoteFeeAmount", 0))) + + await clmm_repo.update_position_fees( + position_address=position.position_address, + base_fee_pending=base_fee_pending, + quote_fee_pending=quote_fee_pending + ) + + logger.debug(f"Refreshed position {position.position_address}: price={current_price}, in_range={in_range}, " + f"base={base_token_amount}, quote={quote_token_amount}") + + except Exception as e: + logger.error(f"Error refreshing position {position.position_address}: {e}", exc_info=True) + raise + + +class GatewayCLMMService(RepositoryService): + """Position and event persistence for the CLMM routes.""" + + repository_class = GatewayCLMMRepository + + # ------------------------------------------------------------------ + # Reads + # ------------------------------------------------------------------ + + async def get_position_wallet(self, position_address: str) -> Optional[str]: + """Wallet recorded for a position, or None if hapi has no row for it. + + Used by close/collect to resolve the signer: an explicit request value wins, + then this, then the chain's default wallet. + """ + async def _fn(clmm_repo): + position = await clmm_repo.get_position_by_address(position_address) + return position.wallet_address if position else None + + return await self._in_repo(_fn) + + async def get_position_pool_address(self, position_address: str) -> Optional[str]: + """Pool a position sits in, or None if hapi has no row for it.""" + async def _fn(clmm_repo): + position = await clmm_repo.get_position_by_address(position_address) + return position.pool_address if position else None + + return await self._in_repo(_fn) + + async def get_position_events( + self, + position_address: str, + event_type: Optional[str] = None, + limit: int = 100, + ) -> Dict[str, Any]: + """Event history for a position, as the endpoint's envelope.""" + async def _fn(clmm_repo): + events = await clmm_repo.get_position_events( + position_address=position_address, + event_type=event_type, + limit=limit + ) + + return { + "data": [clmm_repo.event_to_dict(event) for event in events], + "total_count": len(events) + } + + return await self._in_repo(_fn) + + async def search_positions( + self, + *, + network: Optional[str] = None, + connector: Optional[str] = None, + wallet_address: Optional[str] = None, + trading_pair: Optional[str] = None, + status: Optional[str] = None, + position_addresses: Optional[List[str]] = None, + limit: int = 50, + offset: int = 0, + refresh: bool = False, + gateway_client=None, + ) -> Dict[str, Any]: + """Search stored positions, optionally refreshing them from Gateway first. + + Args: + refresh: Re-read each matched position from Gateway and write back what + it says before answering. Requires ``gateway_client``. + """ + # Validate limit + if limit > 1000: + limit = 1000 + + filters = dict( + network=network, + connector=connector, + wallet_address=wallet_address, + trading_pair=trading_pair, + status=status, + position_addresses=position_addresses, + limit=limit, + offset=offset, + ) + + # Optionally refresh position data from Gateway first + if refresh and gateway_client is not None and await gateway_client.ping(): + await self._refresh_positions(filters, gateway_client) + + # Get final results after refresh + async def _fn(clmm_repo): + positions = await clmm_repo.get_positions(**filters) + + # Get total count for pagination + has_more = len(positions) == limit + + return { + "data": [clmm_repo.position_to_dict(pos) for pos in positions], + "pagination": { + "limit": limit, + "offset": offset, + "has_more": has_more, + "total_count": len(positions) + offset if not has_more else None + } + } + + return await self._in_repo(_fn) + + async def _refresh_positions(self, filters: Dict[str, Any], gateway_client) -> None: + """Re-read every position matching ``filters`` from Gateway, one session each.""" + # Get positions to refresh + async def _fn(clmm_repo): + positions_to_refresh = await clmm_repo.get_positions(**filters) + + # Extract position addresses and details before closing session + return [ + { + "position_address": pos.position_address, + "pool_address": pos.pool_address, + "connector": pos.connector, + "network": pos.network, + "wallet_address": pos.wallet_address + } + for pos in positions_to_refresh + ] + + position_details = await self._in_repo(_fn) + + # Refresh each position in a separate session + logger.info(f"Refreshing {len(position_details)} positions from Gateway") + for pos_detail in position_details: + async def _refresh(clmm_repo, address=pos_detail["position_address"]): + # Get position again in this session + position = await clmm_repo.get_position_by_address(address) + if position: + await refresh_position_data(position, gateway_client, clmm_repo) + + try: + await self._in_repo(_refresh) + except Exception as e: + logger.warning(f"Failed to refresh position {pos_detail['position_address']}: {e}") + # Continue with other positions even if one fails + + # ------------------------------------------------------------------ + # Writes + # ------------------------------------------------------------------ + + async def record_open_position( + self, + *, + position_address: str, + pool_address: str, + network: str, + connector: str, + wallet_address: str, + trading_pair: str, + base_token: str, + quote_token: str, + lower_price: Decimal, + upper_price: Decimal, + entry_price: Optional[float], + base_amount_added: Any, + quote_amount_added: Any, + position_rent: Any, + transaction_hash: str, + gas_fee: Any, + gas_token: Optional[str], + tx_status: str, + ) -> None: + """Record a newly opened position and its OPEN event.""" + position_data = build_position_row( + position_address=position_address, + pool_address=pool_address, + network=network, + connector=connector, + wallet_address=wallet_address, + trading_pair=trading_pair, + base_token=base_token, + quote_token=quote_token, + lower_price=lower_price, + upper_price=upper_price, + entry_price=entry_price, # Pool price when position opened + current_price=entry_price, # Same as entry at open time, updated by poller + initial_base_token_amount=base_amount_added, + initial_quote_token_amount=quote_amount_added, + base_token_amount=base_amount_added, + quote_token_amount=quote_amount_added, + position_rent=position_rent, + # in_range is left UNKNOWN rather than derived from the entry price: the + # poller re-reads the position from the chain within the minute and is the + # only thing that has ever set this column from an observation. + ) + + async def _fn(clmm_repo): + position = await clmm_repo.create_position(position_data) + logger.info(f"Recorded CLMM position in database: {position_address}") + + # Create OPEN event with polled status + event_data = { + "position_id": position.id, + "transaction_hash": transaction_hash, + "event_type": "OPEN", + "base_token_amount": float(base_amount_added) if base_amount_added is not None else None, + "quote_token_amount": float(quote_amount_added) if quote_amount_added is not None else None, + "gas_fee": float(gas_fee) if gas_fee is not None else None, + "gas_token": gas_token, + "status": tx_status + } + + await clmm_repo.create_event(event_data) + logger.info(f"Recorded CLMM OPEN event in database: {transaction_hash} " + f"(status: {tx_status}, gas: {gas_fee} {gas_token})") + + await self._in_repo_best_effort( + _fn, error_message="Error recording CLMM position in database") + + async def record_add_liquidity( + self, + *, + position_address: str, + transaction_hash: str, + tx_status: str, + base_amount_added: Any, + quote_amount_added: Any, + gas_fee: Any, + gas_token: Optional[str], + add_price: Optional[float], + ) -> None: + """Record an ADD_LIQUIDITY event and book the added capital.""" + async def _fn(clmm_repo): + # Get position to link event + position = await clmm_repo.get_position_by_address(position_address) + if not position: + logger.warning(f"ADD_LIQUIDITY {transaction_hash} executed for position " + f"{position_address} with no database record — " + "no event recorded (position may be a pending open " + "not yet discovered)") + return + + event_data = { + "position_id": position.id, + "transaction_hash": transaction_hash, + "event_type": "ADD_LIQUIDITY", + # `is not None`: 0 is a real amount on single-sided adds + "base_token_amount": float(base_amount_added) if base_amount_added is not None else None, + "quote_token_amount": float(quote_amount_added) if quote_amount_added is not None else None, + "gas_fee": float(gas_fee) if gas_fee is not None else None, + "gas_token": gas_token, + "status": tx_status + } + await clmm_repo.create_event(event_data) + logger.info(f"Recorded CLMM ADD_LIQUIDITY event: {transaction_hash} " + f"(status: {tx_status}, gas: {gas_fee} {gas_token})") + + # Added capital raises both the PnL baseline and the held amounts. + # Book here only when the tx confirmed inline (the event is created + # CONFIRMED and the poller never re-processes it); SUBMITTED events + # are booked by the poller's confirm path. + if tx_status == "CONFIRMED": + await clmm_repo.add_to_position_amounts( + position_address=position_address, + base_delta=Decimal(str(base_amount_added or 0)), + quote_delta=Decimal(str(quote_amount_added or 0)), + entry_price=Decimal(str(add_price)) if add_price else None, + ) + + await self._in_repo_best_effort( + _fn, error_message="Error recording ADD_LIQUIDITY event") + + async def record_remove_liquidity( + self, + *, + position_address: str, + transaction_hash: str, + tx_status: str, + base_amount_removed: Any, + quote_amount_removed: Any, + gas_fee: Any, + gas_token: Optional[str], + ) -> None: + """Record a REMOVE_LIQUIDITY event and book the withdrawn capital.""" + async def _fn(clmm_repo): + # Get position to link event + position = await clmm_repo.get_position_by_address(position_address) + if not position: + logger.warning(f"REMOVE_LIQUIDITY {transaction_hash} executed for position " + f"{position_address} with no database record — " + "no event recorded (position may be a pending open " + "not yet discovered)") + return + + # No "percentage" key: GatewayCLMMEvent has no such column, and the + # stray kwarg made create_event raise — silently losing every + # REMOVE_LIQUIDITY event to the log-and-continue handler below. + event_data = { + "position_id": position.id, + "transaction_hash": transaction_hash, + "event_type": "REMOVE_LIQUIDITY", + "base_token_amount": float(base_amount_removed) if base_amount_removed is not None else None, + "quote_token_amount": float(quote_amount_removed) if quote_amount_removed is not None else None, + "gas_fee": float(gas_fee) if gas_fee is not None else None, + "gas_token": gas_token, + "status": tx_status + } + await clmm_repo.create_event(event_data) + logger.info(f"Recorded CLMM REMOVE_LIQUIDITY event: {transaction_hash} " + f"(status: {tx_status}, gas: {gas_fee} {gas_token})") + + # Withdrawn capital lowers both the held amounts and the PnL + # baseline. Book only on inline confirmation (the event is created + # CONFIRMED and the poller never re-processes it); SUBMITTED events + # are booked by the poller's confirm path. + if tx_status == "CONFIRMED": + await clmm_repo.subtract_from_position_amounts( + position_address=position_address, + base_delta=Decimal(str(base_amount_removed or 0)), + quote_delta=Decimal(str(quote_amount_removed or 0)), + ) + + await self._in_repo_best_effort( + _fn, error_message="Error recording REMOVE_LIQUIDITY event") + + async def record_close( + self, + *, + position_address: str, + connector: str, + network: str, + transaction_hash: str, + tx_status: str, + base_amount_removed: Any, + quote_amount_removed: Any, + base_fee_collected: Decimal, + quote_fee_collected: Decimal, + position_rent_refunded: Any, + close_price: Optional[float], + gas_fee: Any, + gas_token: Optional[str], + gateway_client, + ) -> None: + """Record a CLOSE event, book the final fees, and mark the position closed. + + ``gateway_client`` is re-read after the write to confirm the position really + is gone before the row is marked CLOSED. + """ + async def _fn(clmm_repo): + # Get position to link event + position = await clmm_repo.get_position_by_address(position_address) + if not position: + # H8 window: a close on a position hapi has no row for (e.g. a + # pending open awaiting the discovery sweep) leaves no event — + # say so loudly instead of silently skipping. + logger.warning(f"CLOSE {transaction_hash} executed for position " + f"{position_address} with no database record — " + "no CLOSE event recorded (position may be a pending open " + "not yet discovered)") + return + + # Create event record + event_data = { + "position_id": position.id, + "transaction_hash": transaction_hash, + "event_type": "CLOSE", + "base_token_amount": float(base_amount_removed) if base_amount_removed is not None else None, + "quote_token_amount": float(quote_amount_removed) if quote_amount_removed is not None else None, + "base_fee_collected": float(base_fee_collected) if base_fee_collected is not None else None, + "quote_fee_collected": float(quote_fee_collected) if quote_fee_collected is not None else None, + "gas_fee": float(gas_fee) if gas_fee is not None else None, + "gas_token": gas_token, + "status": tx_status + } + await clmm_repo.create_event(event_data) + logger.info(f"Recorded CLMM CLOSE event: {transaction_hash} " + f"(status: {tx_status}, gas: {gas_fee} {gas_token})") + + # Position bookkeeping happens exactly once, when the tx is known + # good: CONFIRMED here (the event is created CONFIRMED, so the + # poller never touches it), or in the poller's confirm path for + # SUBMITTED events. A FAILED tx mutates nothing — the old + # unconditional booking permanently inflated *_fee_collected on + # failed closes. + if tx_status != "CONFIRMED": + return False + + new_base_collected = Decimal(str(position.base_fee_collected)) + base_fee_collected + new_quote_collected = Decimal(str(position.quote_fee_collected)) + quote_fee_collected + + await clmm_repo.update_position_fees( + position_address=position_address, + base_fee_collected=new_base_collected, + quote_fee_collected=new_quote_collected, + base_fee_pending=Decimal("0"), + quote_fee_pending=Decimal("0") + ) + + # Update current_price with close price + if close_price: + await clmm_repo.update_position_liquidity( + position_address=position_address, + base_token_amount=Decimal(str(position.base_token_amount)), + quote_token_amount=Decimal(str(position.quote_token_amount)), + current_price=Decimal(str(close_price)) + ) + + return True + + booked = await self._in_repo_best_effort( + _fn, error_message="Error recording CLOSE event", default=False + ) + if not booked: + return + + # The propagation wait happens with no session held: the writes above are + # committed and the pooled connection is back before we sit idle for two + # seconds, so a fleet closing positions at once cannot drain the pool + # waiting on the chain. + # + # Verify position is actually gone on Gateway before marking CLOSED (some + # connectors 500 instead of 404 for a nonexistent position — right after our + # own close, either means gone). + try: + await asyncio.sleep(2) # Wait for transaction to propagate + + verify_result = await gateway_client.clmm_position_info( + connector=connector, + chain_network=network, + position_address=position_address + ) + + if verify_result and isinstance(verify_result, dict) and "error" in verify_result: + status_code = verify_result.get("status") + if status_code in (404, 500): + async def _close(clmm_repo): + await clmm_repo.close_position( + position_address, + position_rent_refunded=(Decimal(str(position_rent_refunded)) + if position_rent_refunded is not None else None) + ) + + await self._in_repo(_close) + logger.info(f"Position {position_address} verified as closed " + f"(Gateway returned {status_code})") + else: + logger.warning(f"Unexpected error verifying position close: {verify_result}") + elif verify_result and "address" in verify_result: + # Position still exists - might be a failed close or delayed propagation + logger.warning(f"Position {position_address} still exists after close " + "transaction. Will be handled by poller.") + else: + logger.debug("Could not verify position close status, will be handled by poller") + + except Exception as verify_error: + logger.warning(f"Error verifying position close: {verify_error}. Will be handled by poller.") + + logger.info(f"Updated position {position_address}: " + "collected fees updated, pending fees reset to 0.") + + async def record_collect_fees( + self, + *, + position_address: str, + transaction_hash: str, + tx_status: str, + base_fee_collected: Decimal, + quote_fee_collected: Decimal, + gas_fee: Any, + gas_token: Optional[str], + ) -> None: + """Record a COLLECT_FEES event and book the collected fees.""" + async def _fn(clmm_repo): + # Get position to link event + position = await clmm_repo.get_position_by_address(position_address) + if not position: + logger.warning(f"COLLECT_FEES {transaction_hash} executed for position " + f"{position_address} with no database record — " + "no event recorded (position may be a pending open " + "not yet discovered)") + return + + # Create event record + event_data = { + "position_id": position.id, + "transaction_hash": transaction_hash, + "event_type": "COLLECT_FEES", + "base_fee_collected": float(base_fee_collected) if base_fee_collected is not None else None, + "quote_fee_collected": float(quote_fee_collected) if quote_fee_collected is not None else None, + "gas_fee": float(gas_fee) if gas_fee is not None else None, + "gas_token": gas_token, + "status": tx_status + } + await clmm_repo.create_event(event_data) + logger.info(f"Recorded CLMM COLLECT_FEES event: {transaction_hash} " + f"(status: {tx_status}, gas: {gas_fee} {gas_token})") + + # Book fees exactly once: CONFIRMED here (event created CONFIRMED, + # never re-processed), SUBMITTED in the poller's confirm path. + # The old unconditional booking double-counted every pending + # collect (endpoint + poller) and kept phantom fees on failures. + if tx_status == "CONFIRMED": + new_base_collected = Decimal(str(position.base_fee_collected)) + base_fee_collected + new_quote_collected = Decimal(str(position.quote_fee_collected)) + quote_fee_collected + + await clmm_repo.update_position_fees( + position_address=position_address, + base_fee_collected=new_base_collected, + quote_fee_collected=new_quote_collected, + base_fee_pending=Decimal("0"), + quote_fee_pending=Decimal("0") + ) + logger.info(f"Updated position {position_address}: " + "collected fees updated, pending fees reset to 0") + + await self._in_repo_best_effort( + _fn, error_message="Error recording COLLECT_FEES event") + + async def record_failed_write( + self, + *, + transaction_hash: Optional[str], + error: Exception, + event_type: str, + position_address: Optional[str], + ) -> None: + """Record a write that reached the chain and reverted, before the error is re-raised. + + The recording methods above only run when Gateway *returns*. A transaction that + landed and reverted does not return: Gateway raises, the client turns it into a + GatewayError, and control skips every ``create_event`` call to land in an + ``except`` that persists nothing. So the database said every operation ever + attempted had succeeded, while a close that reverted at slot 440494812 — costing + 0.000011772 SOL — left no row at all. + + Only failures carrying a transaction id are recorded (the caller parses it out of + the error). A pre-flight simulation failure never got one and cost nothing, and + inventing an identifier for it would put a row in the table that no lookup by hash + could ever match. + + Recording never masks the original failure: the caller still gets Gateway's error. + """ + if not transaction_hash or not position_address: + return + + async def _fn(repo): + position = await repo.get_position_by_address(position_address) + if position is None: + logger.warning( + f"CLMM {event_type} {transaction_hash} reverted on-chain for position " + f"{position_address}, which has no database record — no event written." + ) + return False + await repo.create_event({ + "position_id": position.id, + "transaction_hash": transaction_hash, + "event_type": event_type, + "status": "FAILED", + "error_message": str(error), + }) + return True + + recorded = await self._in_repo_best_effort( + _fn, error_message=f"Error recording failed CLMM {event_type}", default=False) + if recorded: + logger.error( + f"CLMM {event_type} {transaction_hash} landed on-chain and FAILED for position " + f"{position_address}; recorded. {error}" + ) + + # ------------------------------------------------------------------ + # Transaction poller + # + # The poller used to construct GatewayCLMMRepository itself in five places and + # hold one session open across a whole poll cycle's worth of Gateway calls. It + # now decides *what the chain says* — status classification, the dropped grace + # window, the consecutive-miss gate — and everything about *what gets written* + # lives here with the routes' writes. Each operation takes its own short session, + # so a failure part-way through a cycle no longer discards the statuses already + # confirmed in it. + # ------------------------------------------------------------------ + + async def get_pending_events(self, limit: int = 100) -> List[Dict[str, Any]]: + """Events still awaiting confirmation, with the network their position sits on. + + Returned as plain dicts: the poller does Gateway I/O between reading these and + writing the result, and ORM instances must not outlive their session. + ``network`` is None when the event's position row is missing, which the caller + reports rather than guessing a chain from. + """ + async def _fn(clmm_repo): + events = await clmm_repo.get_pending_events(limit=limit) + pending = [] + for event in events: + position = await clmm_repo.get_position_by_id(event.position_id) + pending.append({ + "transaction_hash": event.transaction_hash, + "timestamp": event.timestamp, + "network": position.network if position else None, + "position_address": position.position_address if position else None, + }) + return pending + + return await self._in_repo(_fn) + + async def update_event_status( + self, + *, + transaction_hash: str, + status: str, + error_message: Optional[str] = None, + gas_fee: Optional[Decimal] = None, + gas_token: Optional[str] = None, + ) -> None: + """Record what a poll found for an event that did not confirm.""" + async def _fn(clmm_repo): + await clmm_repo.update_event_status( + transaction_hash=transaction_hash, + status=status, + error_message=error_message, + gas_fee=gas_fee, + gas_token=gas_token, + ) + + await self._in_repo_best_effort( + _fn, error_message=f"Error recording {status} status for CLMM event {transaction_hash}") + + async def record_event_confirmed( + self, + *, + transaction_hash: str, + gas_fee: Optional[Decimal] = None, + gas_token: Optional[str] = None, + ) -> None: + """Mark an event CONFIRMED and apply the bookkeeping it owes its position. + + Both halves share one session, as they always have. The bookkeeping is guarded + separately: a position that cannot be booked must not roll back the confirmed + status, or the next cycle would poll the same transaction and book it twice. + """ + async def _fn(clmm_repo): + event = await clmm_repo.update_event_status( + transaction_hash=transaction_hash, + status="CONFIRMED", + gas_fee=gas_fee, + gas_token=gas_token, + ) + if event is None: + logger.warning(f"CLMM event {transaction_hash} confirmed on-chain but has no row to update") + return + + try: + await self._book_confirmed_event(clmm_repo, event) + except Exception as e: + logger.error(f"Error updating position from event {event.id}: {e}", exc_info=True) + + await self._in_repo_best_effort( + _fn, error_message=f"Error confirming CLMM event {transaction_hash}") + + @staticmethod + async def _book_confirmed_event(clmm_repo, event) -> None: + """Apply a newly confirmed event's effect to its position. + + Fee and capital booking happens exactly once, here: the routes only mutate the + position when Gateway confirmed the transaction inline (those events are created + CONFIRMED and never reach this path), and leave submitted-not-confirmed booking + to the poller. + """ + position = await clmm_repo.get_position_by_id(event.position_id) + if not position: + logger.error(f"Position not found for event {event.id}") + return + + if event.event_type == "CLOSE": + if event.base_fee_collected is not None or event.quote_fee_collected is not None: + new_base = float(position.base_fee_collected or 0) + float(event.base_fee_collected or 0) + new_quote = float(position.quote_fee_collected or 0) + float(event.quote_fee_collected or 0) + await clmm_repo.update_position_fees( + position_address=position.position_address, + base_fee_collected=Decimal(str(new_base)), + quote_fee_collected=Decimal(str(new_quote)), + base_fee_pending=Decimal("0"), + quote_fee_pending=Decimal("0") + ) + await clmm_repo.close_position(position.position_address) + + elif event.event_type == "ADD_LIQUIDITY": + # Added capital raises both the PnL baseline and the held amounts. Event + # amounts may be the requested figures (recorded at submit time) rather + # than on-chain actuals — the accepted residual is that pending-tx amounts + # are not backfilled from txData; requested amounts are the best available. + if event.base_token_amount or event.quote_token_amount: + await clmm_repo.add_to_position_amounts( + position_address=position.position_address, + base_delta=Decimal(str(event.base_token_amount or 0)), + quote_delta=Decimal(str(event.quote_token_amount or 0)), + ) + + elif event.event_type == "REMOVE_LIQUIDITY": + # The mirror of ADD_LIQUIDITY: withdrawn capital lowers both the held + # amounts and the PnL baseline. + if event.base_token_amount or event.quote_token_amount: + await clmm_repo.subtract_from_position_amounts( + position_address=position.position_address, + base_delta=Decimal(str(event.base_token_amount or 0)), + quote_delta=Decimal(str(event.quote_token_amount or 0)), + ) + + elif event.event_type == "COLLECT_FEES": + if event.base_fee_collected or event.quote_fee_collected: + new_base_collected = float(position.base_fee_collected or 0) + float(event.base_fee_collected or 0) + new_quote_collected = float(position.quote_fee_collected or 0) + float(event.quote_fee_collected or 0) + await clmm_repo.update_position_fees( + position_address=position.position_address, + base_fee_collected=Decimal(str(new_base_collected)), + quote_fee_collected=Decimal(str(new_quote_collected)), + base_fee_pending=Decimal("0"), + quote_fee_pending=Decimal("0") + ) + + async def get_tracked_position_addresses(self, recently_closed_seconds: int) -> Dict[str, Set[str]]: + """The three address sets the discovery sweep compares Gateway's listing against. + + One session for all three: they are read together and used together, and a + position that changed status between them would make the sweep contradict itself. + """ + async def _fn(clmm_repo): + return { + "open": await clmm_repo.get_position_addresses_set(status="OPEN"), + "closed": await clmm_repo.get_position_addresses_set(status="CLOSED"), + "recently_closed": await clmm_repo.get_recently_closed_addresses(recently_closed_seconds), + } + + return await self._in_repo(_fn) + + async def reopen_position(self, position_address: str) -> bool: + """Undo a close for a position the chain still reports as live. True if reopened.""" + async def _fn(clmm_repo): + return await clmm_repo.reopen_position(position_address) is not None + + return await self._in_repo_best_effort( + _fn, error_message=f"Error reopening position {position_address}", default=False) + + async def record_discovered_position( + self, + *, + pos_data: Dict[str, Any], + connector: str, + network: str, + wallet_address: str, + ) -> bool: + """Record a position the poller found on-chain that hapi has no row for. + + These were opened elsewhere (the UI, an executor talking to Gateway directly), + so the entry price and the initial deposit are unknowable: what the chain holds + right now is the best available estimate for both, and is recorded as such. + + The row itself is :func:`build_position_row`'s — the same columns the open route + writes — and it is paired with a synthetic DISCOVERED event so the history says + where the row came from. + """ + position_address = pos_data.get("address") + if not position_address: + return False + + # Full token addresses are used as the token identity here, as the open route does. + base_token = pos_data.get("baseTokenAddress") or "UNKNOWN" + quote_token = pos_data.get("quoteTokenAddress") or "UNKNOWN" + + current_price = float(pos_data.get("price", 0)) + lower_price = float(pos_data.get("lowerPrice", 0)) + upper_price = float(pos_data.get("upperPrice", 0)) + + base_token_amount = float(pos_data.get("baseTokenAmount", 0)) + quote_token_amount = float(pos_data.get("quoteTokenAmount", 0)) + + in_range = "UNKNOWN" + if current_price > 0 and lower_price > 0 and upper_price > 0: + in_range = "IN_RANGE" if lower_price <= current_price <= upper_price else "OUT_OF_RANGE" + + position_data = build_position_row( + position_address=position_address, + pool_address=pos_data.get("poolAddress", ""), + network=network, + connector=connector, + wallet_address=wallet_address, + trading_pair=f"{base_token}-{quote_token}", + base_token=base_token, + quote_token=quote_token, + lower_price=lower_price, + upper_price=upper_price, + lower_bin_id=pos_data.get("lowerBinId"), + upper_bin_id=pos_data.get("upperBinId"), + entry_price=current_price, # Best available estimate + current_price=current_price, + # Nothing records what was originally deposited, so what is held now stands + # in for it — the PnL baseline starts at the moment of discovery. + initial_base_token_amount=base_token_amount, + initial_quote_token_amount=quote_token_amount, + base_token_amount=base_token_amount, + quote_token_amount=quote_token_amount, + in_range=in_range, + base_fee_pending=float(pos_data.get("baseFeeAmount", 0)), + quote_fee_pending=float(pos_data.get("quoteFeeAmount", 0)), + ) + + async def _fn(clmm_repo): + position = await clmm_repo.create_position(position_data) + await clmm_repo.create_event({ + "position_id": position.id, + # Synthetic: there is no transaction of ours to confirm. + "transaction_hash": f"discovered_{position_address[:16]}", + "event_type": "DISCOVERED", + "base_token_amount": base_token_amount, + "quote_token_amount": quote_token_amount, + "status": "CONFIRMED", + }) + return True + + return await self._in_repo_best_effort( + _fn, error_message=f"Error creating discovered position {position_address}", default=False) + + async def get_open_positions(self) -> List[Dict[str, Any]]: + """The open positions the poller refreshes, as plain dicts. + + Only the fields needed to re-read a position from Gateway: the refresh does + network I/O per position and must not hold a session while it does. + """ + async def _fn(clmm_repo): + positions = await clmm_repo.get_open_positions() + return [ + { + "id": position.id, + "position_address": position.position_address, + "wallet_address": position.wallet_address, + "connector": position.connector, + "network": position.network, + } + for position in positions + ] + + return await self._in_repo(_fn) + + async def mark_position_closed(self, position_address: str) -> None: + """Close a position the chain no longer reports. Never fails the poll cycle.""" + async def _fn(clmm_repo): + await clmm_repo.close_position(position_address) + + await self._in_repo_best_effort( + _fn, error_message=f"Error closing position {position_address}") + + async def record_position_state( + self, + *, + position_address: str, + base_token_amount: Decimal, + quote_token_amount: Decimal, + in_range: str, + current_price: Decimal, + base_fee_pending: Decimal, + quote_fee_pending: Decimal, + ) -> None: + """Write back everything one on-chain read of a position says, in one session. + + Pending fees are always written: 0 is a real value right after an external + collect, and a non-zero guard would leave the old figure standing forever. + """ + async def _fn(clmm_repo): + await clmm_repo.update_position_liquidity( + position_address=position_address, + base_token_amount=base_token_amount, + quote_token_amount=quote_token_amount, + in_range=in_range, + current_price=current_price, + ) + await clmm_repo.update_position_fees( + position_address=position_address, + base_fee_pending=base_fee_pending, + quote_fee_pending=quote_fee_pending, + ) + + await self._in_repo_best_effort( + _fn, error_message=f"Error updating state for position {position_address}") diff --git a/services/gateway_swap_service.py b/services/gateway_swap_service.py new file mode 100644 index 00000000..93ceb8ac --- /dev/null +++ b/services/gateway_swap_service.py @@ -0,0 +1,194 @@ +"""Persistence for DEX swaps executed through the ``/gateway/swap*`` routes. + +The shape of a ``gateway_swaps`` row, the pagination envelope over a swap search and +the "record it but never fail the swap over it" policy all live here rather than in +the handlers, which only decide what to answer the caller. +""" +import logging +from decimal import Decimal +from typing import Any, Dict, List, Optional + +from database.repositories import GatewaySwapRepository +from services.repository_service import RepositoryService + +logger = logging.getLogger(__name__) + + +class GatewaySwapService(RepositoryService): + """Swap history persistence for the swap routes.""" + + repository_class = GatewaySwapRepository + + async def record_swap( + self, + *, + transaction_hash: str, + network: str, + connector: str, + wallet_address: str, + trading_pair: str, + base_token: str, + quote_token: str, + side: str, + input_amount: Decimal, + output_amount: Decimal, + price: Decimal, + slippage_pct: Optional[Decimal], + gas_fee: Optional[Decimal], + gas_token: Optional[str], + status: str, + pool_address: Optional[str], + ) -> None: + """Record a settled swap. Best-effort: the swap already happened on-chain.""" + async def _fn(swap_repo): + swap_data = { + "transaction_hash": transaction_hash, + "network": network, + # Store the base venue name: a swap on "jupiter" and one on + # "jupiter/router" are the same venue and must file together. + "connector": connector.split("/")[0], + "wallet_address": wallet_address, + "trading_pair": trading_pair, + "base_token": base_token, + "quote_token": quote_token, + "side": side, + "input_amount": float(input_amount), + "output_amount": float(output_amount), + "price": float(price), + "slippage_pct": float(slippage_pct) if slippage_pct is not None else None, + "gas_fee": float(gas_fee) if gas_fee is not None else None, + "gas_token": gas_token, + "status": status, + # Set by the pool-scoped routes, which resolve exactly one pool; a + # router picks its own path across pools and leaves it unset. + "pool_address": pool_address + } + + await swap_repo.create_swap(swap_data) + logger.info(f"Recorded swap in database: {transaction_hash} (status: {status})") + + await self._in_repo_best_effort( + _fn, error_message="Error recording swap in database") + + async def get_pending_swaps(self, limit: int = 100) -> List[Dict[str, Any]]: + """Swaps still awaiting confirmation, as plain dicts. + + The poller does Gateway I/O between reading these and writing the result, so + the session is closed before it starts and ORM instances must not outlive it. + """ + async def _fn(swap_repo): + swaps = await swap_repo.get_pending_swaps(limit=limit) + return [ + { + "transaction_hash": swap.transaction_hash, + "network": swap.network, + "timestamp": swap.timestamp, + } + for swap in swaps + ] + + return await self._in_repo(_fn) + + async def update_swap_status( + self, + *, + transaction_hash: str, + status: str, + error_message: Optional[str] = None, + gas_fee: Optional[Decimal] = None, + gas_token: Optional[str] = None, + ) -> None: + """Record what polling the chain found for a submitted swap. + + Best-effort like every other write here: the swap settled (or did not) on-chain + regardless, and a database problem must not stop the poll cycle reaching the + rest of the pending book. + """ + async def _fn(swap_repo): + await swap_repo.update_swap_status( + transaction_hash=transaction_hash, + status=status, + error_message=error_message, + gas_fee=gas_fee, + gas_token=gas_token, + ) + + await self._in_repo_best_effort( + _fn, error_message=f"Error recording {status} status for swap {transaction_hash}") + + async def get_swap(self, transaction_hash: str) -> Optional[Dict[str, Any]]: + """One swap by transaction hash, or None when there is no such row. + + None rather than an exception: the caller turns "no row" into its own 404, + which a swallowed database error must never be mistaken for. + """ + async def _fn(swap_repo): + swap = await swap_repo.get_swap_by_tx_hash(transaction_hash) + return swap_repo.to_dict(swap) if swap else None + + return await self._in_repo(_fn) + + async def search_swaps( + self, + *, + network: Optional[str] = None, + connector: Optional[str] = None, + wallet_address: Optional[str] = None, + trading_pair: Optional[str] = None, + status: Optional[str] = None, + start_time: Optional[int] = None, + end_time: Optional[int] = None, + limit: int = 50, + offset: int = 0, + ) -> Dict[str, Any]: + """Filtered swap history, as the endpoint's paginated envelope.""" + # Validate limit + if limit > 1000: + limit = 1000 + + async def _fn(swap_repo): + swaps = await swap_repo.get_swaps( + network=network, + connector=connector, + wallet_address=wallet_address, + trading_pair=trading_pair, + status=status, + start_time=start_time, + end_time=end_time, + limit=limit, + offset=offset + ) + + # Get total count for pagination (simplified - actual count would need separate query) + has_more = len(swaps) == limit + + return { + "data": [swap_repo.to_dict(swap) for swap in swaps], + "pagination": { + "limit": limit, + "offset": offset, + "has_more": has_more, + "total_count": len(swaps) + offset if not has_more else None + } + } + + return await self._in_repo(_fn) + + async def get_swaps_summary( + self, + *, + network: Optional[str] = None, + wallet_address: Optional[str] = None, + start_time: Optional[int] = None, + end_time: Optional[int] = None, + ) -> Dict[str, Any]: + """Aggregate swap statistics (volume, fees, success rate).""" + async def _fn(swap_repo): + return await swap_repo.get_swaps_summary( + network=network, + wallet_address=wallet_address, + start_time=start_time, + end_time=end_time + ) + + return await self._in_repo(_fn) diff --git a/services/gateway_transaction_poller.py b/services/gateway_transaction_poller.py index bca5a0c3..fdce6c29 100644 --- a/services/gateway_transaction_poller.py +++ b/services/gateway_transaction_poller.py @@ -13,9 +13,9 @@ from typing import Dict, Optional from database import AsyncDatabaseManager -from database.models import GatewayCLMMPosition -from database.repositories import GatewayCLMMRepository, GatewaySwapRepository from services.gateway_client import GatewayClient, get_native_gas_token +from services.gateway_clmm_service import GatewayCLMMService +from services.gateway_swap_service import GatewaySwapService logger = logging.getLogger(__name__) @@ -41,6 +41,13 @@ def __init__( ): self.db_manager = db_manager self.gateway_client = gateway_client + # Every row this poller reads or writes goes through the same services the + # /gateway/clmm and /gateway/swap routes persist through, so a position it + # discovers has the same shape as one a route opened. Constructed here rather + # than injected: a RepositoryService holds nothing but the shared db_manager, + # so there is no instance worth threading down from main.py. + self.clmm_service = GatewayCLMMService(db_manager=db_manager) + self.swap_service = GatewaySwapService(db_manager=db_manager) self.poll_interval = poll_interval self.position_poll_interval = position_poll_interval self.max_retry_age = max_retry_age @@ -125,34 +132,32 @@ async def _poll_pending_transactions(self): logger.warning("Gateway not available; skipping transaction poll cycle") return - async with self.db_manager.get_session_context() as session: - swap_repo = GatewaySwapRepository(session) - clmm_repo = GatewayCLMMRepository(session) + # The pending book is read up front and the session released: what follows + # is one Gateway call per transaction, and holding a pooled connection + # across all of them is what used to drain the pool under a busy fleet. + pending_swaps = await self.swap_service.get_pending_swaps(limit=100) + logger.debug(f"Found {len(pending_swaps)} pending swaps") - # Get pending swaps - pending_swaps = await swap_repo.get_pending_swaps(limit=100) - logger.debug(f"Found {len(pending_swaps)} pending swaps") + for swap in pending_swaps: + await self._poll_swap_transaction(swap) - for swap in pending_swaps: - await self._poll_swap_transaction(swap, swap_repo) + pending_events = await self.clmm_service.get_pending_events(limit=100) + logger.debug(f"Found {len(pending_events)} pending CLMM events") - # Get pending CLMM events - pending_events = await clmm_repo.get_pending_events(limit=100) - logger.debug(f"Found {len(pending_events)} pending CLMM events") - - for event in pending_events: - await self._poll_clmm_event_transaction(event, clmm_repo) + for event in pending_events: + await self._poll_clmm_event_transaction(event) except Exception as e: logger.error(f"Error polling pending transactions: {e}", exc_info=True) - async def _poll_swap_transaction(self, swap, swap_repo: GatewaySwapRepository): + async def _poll_swap_transaction(self, swap: Dict): """Poll a specific swap transaction status.""" + transaction_hash = swap["transaction_hash"] try: # Parse network into chain and network - parts = swap.network.split('-', 1) + parts = swap["network"].split('-', 1) if len(parts) != 2: - logger.error(f"Invalid network format for swap {swap.transaction_hash}: {swap.network}") + logger.error(f"Invalid network format for swap {transaction_hash}: {swap['network']}") return chain, network = parts @@ -160,13 +165,13 @@ async def _poll_swap_transaction(self, swap, swap_repo: GatewaySwapRepository): status_result = await self._check_transaction_status( chain=chain, network=network, - tx_hash=swap.transaction_hash + tx_hash=transaction_hash ) if status_result is None: # Transient (Gateway/RPC hiccup): no information, no state change. return - age = (datetime.now(timezone.utc) - swap.timestamp).total_seconds() + age = (datetime.now(timezone.utc) - swap["timestamp"]).total_seconds() status = status_result["status"] gas_fee_raw = status_result.get("gas_fee") gas_fee = Decimal(str(gas_fee_raw)) if gas_fee_raw is not None else None @@ -176,58 +181,58 @@ async def _poll_swap_transaction(self, swap, swap_repo: GatewaySwapRepository): # txData — a swap recorded while pending keeps its request-side leg # and 0 placeholders after confirmation. Backfilling requires parsing # balance changes from txData (deferred by design). - logger.info(f"Swap transaction confirmed: {swap.transaction_hash}") - await swap_repo.update_swap_status( - transaction_hash=swap.transaction_hash, + logger.info(f"Swap transaction confirmed: {transaction_hash}") + await self.swap_service.update_swap_status( + transaction_hash=transaction_hash, status="CONFIRMED", gas_fee=gas_fee, gas_token=status_result.get("gas_token") ) elif status == "FAILED": # A landed-but-failed tx still paid gas — record it. - logger.warning(f"Swap transaction failed: {swap.transaction_hash}") - await swap_repo.update_swap_status( - transaction_hash=swap.transaction_hash, + logger.warning(f"Swap transaction failed: {transaction_hash}") + await self.swap_service.update_swap_status( + transaction_hash=transaction_hash, status="FAILED", error_message=status_result.get("error_message", "Transaction failed on-chain"), gas_fee=gas_fee, gas_token=status_result.get("gas_token") ) elif status == "DROPPED" and age > self.DROPPED_GRACE_SECONDS: - logger.warning(f"Swap transaction dropped (not found on-chain): {swap.transaction_hash}") - await swap_repo.update_swap_status( - transaction_hash=swap.transaction_hash, + logger.warning(f"Swap transaction dropped (not found on-chain): {transaction_hash}") + await self.swap_service.update_swap_status( + transaction_hash=transaction_hash, status="FAILED", error_message="Transaction not found on-chain (dropped after blockhash expiry)" ) elif status == "PENDING" and age > self.max_retry_age: # Genuinely still unconfirmed after a successful poll — only now may # the age timeout fire. - logger.warning(f"Swap {swap.transaction_hash} exceeded max retry age, marking as FAILED") - await swap_repo.update_swap_status( - transaction_hash=swap.transaction_hash, + logger.warning(f"Swap {transaction_hash} exceeded max retry age, marking as FAILED") + await self.swap_service.update_swap_status( + transaction_hash=transaction_hash, status="FAILED", error_message="Transaction confirmation timeout" ) # PENDING within age / DROPPED within grace: retry next cycle. except Exception as e: - logger.error(f"Error polling swap transaction {swap.transaction_hash}: {e}") + logger.error(f"Error polling swap transaction {transaction_hash}: {e}") - async def _poll_clmm_event_transaction(self, event, clmm_repo: GatewayCLMMRepository): + async def _poll_clmm_event_transaction(self, event: Dict): """Poll a specific CLMM event transaction status.""" + transaction_hash = event["transaction_hash"] try: - # Get the position by ID from the event's position_id foreign key - position = await clmm_repo.get_position_by_id(event.position_id) - - if not position: - logger.error(f"Position not found for CLMM event {event.transaction_hash}") + # The network comes from the position the event belongs to; None means + # that row is missing, so there is no chain to ask about this event. + if not event["network"]: + logger.error(f"Position not found for CLMM event {transaction_hash}") return # Parse network - parts = position.network.split('-', 1) + parts = event["network"].split('-', 1) if len(parts) != 2: - logger.error(f"Invalid network format for CLMM event {event.transaction_hash}: {position.network}") + logger.error(f"Invalid network format for CLMM event {transaction_hash}: {event['network']}") return chain, network = parts @@ -235,124 +240,53 @@ async def _poll_clmm_event_transaction(self, event, clmm_repo: GatewayCLMMReposi status_result = await self._check_transaction_status( chain=chain, network=network, - tx_hash=event.transaction_hash + tx_hash=transaction_hash ) if status_result is None: # Transient (Gateway/RPC hiccup): no information, no state change. return - age = (datetime.now(timezone.utc) - event.timestamp).total_seconds() + age = (datetime.now(timezone.utc) - event["timestamp"]).total_seconds() status = status_result["status"] gas_fee_raw = status_result.get("gas_fee") gas_fee = Decimal(str(gas_fee_raw)) if gas_fee_raw is not None else None if status == "CONFIRMED": - logger.info(f"CLMM event transaction confirmed: {event.transaction_hash}") - await clmm_repo.update_event_status( - transaction_hash=event.transaction_hash, - status="CONFIRMED", + logger.info(f"CLMM event transaction confirmed: {transaction_hash}") + # Marks the event CONFIRMED and applies what it owes its position + # (fees booked, capital added or withdrawn, a close finalised). + await self.clmm_service.record_event_confirmed( + transaction_hash=transaction_hash, gas_fee=gas_fee, gas_token=status_result.get("gas_token") ) - # Update position state based on event type - await self._update_position_from_event(event, clmm_repo) elif status == "FAILED": - logger.warning(f"CLMM event transaction failed: {event.transaction_hash}") - await clmm_repo.update_event_status( - transaction_hash=event.transaction_hash, + logger.warning(f"CLMM event transaction failed: {transaction_hash}") + await self.clmm_service.update_event_status( + transaction_hash=transaction_hash, status="FAILED", error_message=status_result.get("error_message", "Transaction failed on-chain"), gas_fee=gas_fee, gas_token=status_result.get("gas_token") ) elif status == "DROPPED" and age > self.DROPPED_GRACE_SECONDS: - logger.warning(f"CLMM event transaction dropped (not found on-chain): {event.transaction_hash}") - await clmm_repo.update_event_status( - transaction_hash=event.transaction_hash, + logger.warning(f"CLMM event transaction dropped (not found on-chain): {transaction_hash}") + await self.clmm_service.update_event_status( + transaction_hash=transaction_hash, status="FAILED", error_message="Transaction not found on-chain (dropped after blockhash expiry)" ) elif status == "PENDING" and age > self.max_retry_age: - logger.warning(f"CLMM event {event.transaction_hash} exceeded max retry age, marking as FAILED") - await clmm_repo.update_event_status( - transaction_hash=event.transaction_hash, + logger.warning(f"CLMM event {transaction_hash} exceeded max retry age, marking as FAILED") + await self.clmm_service.update_event_status( + transaction_hash=transaction_hash, status="FAILED", error_message="Transaction confirmation timeout" ) # PENDING within age / DROPPED within grace: retry next cycle. except Exception as e: - logger.error(f"Error polling CLMM event transaction {event.transaction_hash}: {e}") - - async def _update_position_from_event(self, event, clmm_repo: GatewayCLMMRepository): - """Update CLMM position state based on confirmed event.""" - try: - # Get position by ID using the repository - position = await clmm_repo.get_position_by_id(event.position_id) - - if not position: - logger.error(f"Position not found for event {event.id}") - return - - if event.event_type == "CLOSE": - # Fee booking happens exactly once, on confirmation: the endpoints - # only mutate the position when Gateway confirmed the tx inline, and - # leave submitted-not-confirmed booking to this path. - if event.base_fee_collected is not None or event.quote_fee_collected is not None: - new_base = float(position.base_fee_collected or 0) + float(event.base_fee_collected or 0) - new_quote = float(position.quote_fee_collected or 0) + float(event.quote_fee_collected or 0) - await clmm_repo.update_position_fees( - position_address=position.position_address, - base_fee_collected=Decimal(str(new_base)), - quote_fee_collected=Decimal(str(new_quote)), - base_fee_pending=Decimal("0"), - quote_fee_pending=Decimal("0") - ) - await clmm_repo.close_position(position.position_address) - - elif event.event_type == "ADD_LIQUIDITY": - # Added capital raises both the PnL baseline and the held amounts. - # Event amounts may be the requested figures (recorded at submit time) - # rather than on-chain actuals — the accepted residual is that - # pending-tx amounts are not backfilled from txData; requested amounts - # are the best available. - if event.base_token_amount or event.quote_token_amount: - await clmm_repo.add_to_position_amounts( - position_address=position.position_address, - base_delta=Decimal(str(event.base_token_amount or 0)), - quote_delta=Decimal(str(event.quote_token_amount or 0)), - ) - - elif event.event_type == "REMOVE_LIQUIDITY": - # The mirror of ADD_LIQUIDITY: withdrawn capital lowers both the held - # amounts and the PnL baseline. Endpoints book inline only for txs - # Gateway confirmed at submit time — those events are created CONFIRMED - # and never reach this path, so there is no double count. - if event.base_token_amount or event.quote_token_amount: - await clmm_repo.subtract_from_position_amounts( - position_address=position.position_address, - base_delta=Decimal(str(event.base_token_amount or 0)), - quote_delta=Decimal(str(event.quote_token_amount or 0)), - ) - - elif event.event_type == "COLLECT_FEES": - # Add collected fees to cumulative total (endpoints book inline only - # for txs Gateway confirmed at submit time — those events are created - # CONFIRMED and never reach this path, so there is no double count). - if event.base_fee_collected or event.quote_fee_collected: - new_base_collected = float(position.base_fee_collected or 0) + float(event.base_fee_collected or 0) - new_quote_collected = float(position.quote_fee_collected or 0) + float(event.quote_fee_collected or 0) - - await clmm_repo.update_position_fees( - position_address=position.position_address, - base_fee_collected=Decimal(str(new_base_collected)), - quote_fee_collected=Decimal(str(new_quote_collected)), - base_fee_pending=Decimal("0"), - quote_fee_pending=Decimal("0") - ) - - except Exception as e: - logger.error(f"Error updating position from event: {e}", exc_info=True) + logger.error(f"Error polling CLMM event transaction {transaction_hash}: {e}") async def _check_transaction_status( self, @@ -539,15 +473,13 @@ async def _discover_positions_from_gateway(self) -> int: logger.debug("No wallets configured in Gateway, skipping position discovery") return 0 - # Get existing position addresses from database (for quick existence check) - async with self.db_manager.get_session_context() as session: - clmm_repo = GatewayCLMMRepository(session) - # Get OPEN positions (to skip - already tracked correctly) - open_positions = await clmm_repo.get_position_addresses_set(status="OPEN") - # Get CLOSED positions (to potentially reopen if still on-chain) - closed_positions = await clmm_repo.get_position_addresses_set(status="CLOSED") - # Positions closed moments ago are exempt from reopening (lag guard) - recently_closed = await clmm_repo.get_recently_closed_addresses(self.REOPEN_GRACE_SECONDS) + # Existing position addresses (for quick existence checks): OPEN ones are + # already tracked correctly, CLOSED ones may need reopening if still + # on-chain, and ones closed moments ago are exempt from that (lag guard). + tracked = await self.clmm_service.get_tracked_position_addresses(self.REOPEN_GRACE_SECONDS) + open_positions = tracked["open"] + closed_positions = tracked["closed"] + recently_closed = tracked["recently_closed"] # Poll each supported connector/chain/wallet combination for config in self.SUPPORTED_CLMM_CONFIGS: @@ -593,28 +525,25 @@ async def _discover_positions_from_gateway(self) -> int: "skipping reopen (listing may lag the close)") continue # Position exists on-chain but is CLOSED in DB → reopen it - async with self.db_manager.get_session_context() as session: - clmm_repo = GatewayCLMMRepository(session) - reopened = await clmm_repo.reopen_position(position_address) - if reopened: - reopened_count += 1 - # Move from closed to open set for this run - closed_positions.discard(position_address) - open_positions.add(position_address) - logger.warning(f"Reopened position {position_address} - " - f"was CLOSED in DB but still exists on-chain") + if await self.clmm_service.reopen_position(position_address): + reopened_count += 1 + # Move from closed to open set for this run + closed_positions.discard(position_address) + open_positions.add(position_address) + logger.warning(f"Reopened position {position_address} - " + f"was CLOSED in DB but still exists on-chain") continue - # Create new position in database - new_position = await self._create_discovered_position( + # Create new position in database, in the same shape the + # /gateway/clmm/open route writes. + created = await self.clmm_service.record_discovered_position( pos_data=pos_data, connector=connector, - chain=chain, - network=network, + network=chain_network, wallet_address=wallet_address ) - if new_position: + if created: discovered_count += 1 open_positions.add(position_address) logger.info(f"Discovered new position: {position_address} " @@ -632,145 +561,29 @@ async def _discover_positions_from_gateway(self) -> int: return discovered_count + reopened_count - async def _create_discovered_position( - self, - pos_data: Dict, - connector: str, - chain: str, - network: str, - wallet_address: str - ) -> Optional[GatewayCLMMPosition]: - """ - Create a database record for a discovered position. - - These positions were created externally (e.g., via UI) and are being - discovered by the poller. - """ - try: - position_address = pos_data.get("address") - pool_address = pos_data.get("poolAddress", "") - - # Extract token addresses - base_token_address = pos_data.get("baseTokenAddress", "") - quote_token_address = pos_data.get("quoteTokenAddress", "") - - # Use full addresses as tokens (consistent with API-created positions) - base_token = base_token_address if base_token_address else "UNKNOWN" - quote_token = quote_token_address if quote_token_address else "UNKNOWN" - trading_pair = f"{base_token}-{quote_token}" - - # Extract price data - current_price = float(pos_data.get("price", 0)) - lower_price = float(pos_data.get("lowerPrice", 0)) - upper_price = float(pos_data.get("upperPrice", 0)) - - # Extract liquidity amounts - base_token_amount = float(pos_data.get("baseTokenAmount", 0)) - quote_token_amount = float(pos_data.get("quoteTokenAmount", 0)) - - # Extract fee data - base_fee_pending = float(pos_data.get("baseFeeAmount", 0)) - quote_fee_pending = float(pos_data.get("quoteFeeAmount", 0)) - - # Extract bin IDs (for Meteora) - lower_bin_id = pos_data.get("lowerBinId") - upper_bin_id = pos_data.get("upperBinId") - - # Calculate in_range status - in_range = "UNKNOWN" - if current_price > 0 and lower_price > 0 and upper_price > 0: - if lower_price <= current_price <= upper_price: - in_range = "IN_RANGE" - else: - in_range = "OUT_OF_RANGE" - - # Calculate percentage: (upper_price - lower_price) / lower_price - percentage = None - if lower_price > 0: - percentage = (upper_price - lower_price) / lower_price - - # Network in unified format - network_id = f"{chain}-{network}" - - # Create position in database - async with self.db_manager.get_session_context() as session: - clmm_repo = GatewayCLMMRepository(session) - - position_data = { - "position_address": position_address, - "pool_address": pool_address, - "network": network_id, - "connector": connector, - "wallet_address": wallet_address, - "trading_pair": trading_pair, - "base_token": base_token, - "quote_token": quote_token, - "status": "OPEN", - "lower_price": lower_price, - "upper_price": upper_price, - "lower_bin_id": lower_bin_id, - "upper_bin_id": upper_bin_id, - "entry_price": current_price, # Best available estimate - "current_price": current_price, - "percentage": percentage, - # For discovered positions, we don't know initial amounts - # Use current amounts as initial (best estimate) - "initial_base_token_amount": base_token_amount, - "initial_quote_token_amount": quote_token_amount, - "base_token_amount": base_token_amount, - "quote_token_amount": quote_token_amount, - "in_range": in_range, - "base_fee_pending": base_fee_pending, - "quote_fee_pending": quote_fee_pending, - "base_fee_collected": 0, - "quote_fee_collected": 0, - } - - position = await clmm_repo.create_position(position_data) - - # Create a DISCOVERED event to mark this position was auto-discovered - event_data = { - "position_id": position.id, - "transaction_hash": f"discovered_{position_address[:16]}", # Synthetic tx hash - "event_type": "DISCOVERED", - "base_token_amount": base_token_amount, - "quote_token_amount": quote_token_amount, - "status": "CONFIRMED" # No actual transaction to confirm - } - await clmm_repo.create_event(event_data) - - return position - - except Exception as e: - logger.error(f"Error creating discovered position {pos_data.get('address')}: {e}", exc_info=True) - return None - async def _update_all_open_positions(self): """Update state for all open positions from Gateway.""" try: - async with self.db_manager.get_session_context() as session: - clmm_repo = GatewayCLMMRepository(session) - - # Get all open positions - open_positions = await clmm_repo.get_open_positions() - if not open_positions: - logger.debug("No open CLMM positions to update") - return + open_positions = await self.clmm_service.get_open_positions() + if not open_positions: + logger.debug("No open CLMM positions to update") + return - logger.info(f"Updating {len(open_positions)} open CLMM positions") + logger.info(f"Updating {len(open_positions)} open CLMM positions") - # Update each position within the same session - for position in open_positions: - try: - await self._refresh_position_state(position, clmm_repo) - except Exception as e: - logger.warning(f"Failed to update position {position.position_address}: {e}") - continue + # One Gateway read and one short write per position: the list above is + # already detached, so no session is held across the network calls. + for position in open_positions: + try: + await self._refresh_position_state(position) + except Exception as e: + logger.warning(f"Failed to update position {position['position_address']}: {e}") + continue except Exception as e: logger.error(f"Error updating open positions: {e}", exc_info=True) - async def _refresh_position_state(self, position: GatewayCLMMPosition, clmm_repo: GatewayCLMMRepository): + async def _refresh_position_state(self, position: Dict): """ Refresh a single position's state from Gateway. @@ -780,36 +593,37 @@ async def _refresh_position_state(self, position: GatewayCLMMPosition, clmm_repo - pending fees - position status (if closed externally) """ + position_address = position["position_address"] try: # Validate position has required fields - if not position.position_address: - logger.error(f"Position ID {position.id} has no position_address, skipping refresh") + if not position_address: + logger.error(f"Position ID {position['id']} has no position_address, skipping refresh") return - if not position.wallet_address: - logger.error(f"Position {position.position_address} has no wallet_address, skipping refresh") + if not position["wallet_address"]: + logger.error(f"Position {position_address} has no wallet_address, skipping refresh") return - if not position.connector: - logger.error(f"Position {position.position_address} has no connector, skipping refresh") + if not position["connector"]: + logger.error(f"Position {position_address} has no connector, skipping refresh") return - if not position.network: - logger.error(f"Position {position.position_address} has no network, skipping refresh") + if not position["network"]: + logger.error(f"Position {position_address} has no network, skipping refresh") return # Get individual position info from Gateway (includes pending fees) try: result = await self.gateway_client.clmm_position_info( - connector=position.connector, - chain_network=position.network, # position.network is already in 'chain-network' format - position_address=position.position_address + connector=position["connector"], + chain_network=position["network"], # already in 'chain-network' format + position_address=position_address ) # Check for Gateway errors if result is None: - logger.debug(f"Gateway connection error for position {position.position_address}, skipping update") + logger.debug(f"Gateway connection error for position {position_address}, skipping update") return if not isinstance(result, dict): - logger.warning(f"Unexpected response type for position {position.position_address}: {type(result)}") + logger.warning(f"Unexpected response type for position {position_address}: {type(result)}") return # Check if Gateway returned an error response @@ -821,34 +635,34 @@ async def _refresh_position_state(self, position: GatewayCLMMPosition, clmm_repo # counts as a strike, and the position closes after # MISSING_STRIKES_TO_CLOSE consecutive misses, never on one. if status_code in (404, 500): - strikes = self._position_missing_strikes.get(position.position_address, 0) + 1 - self._position_missing_strikes[position.position_address] = strikes + strikes = self._position_missing_strikes.get(position_address, 0) + 1 + self._position_missing_strikes[position_address] = strikes if strikes >= self.MISSING_STRIKES_TO_CLOSE: - logger.info(f"Position {position.position_address} missing from Gateway " + logger.info(f"Position {position_address} missing from Gateway " f"{strikes} consecutive times (last status: {status_code}), " "marking as CLOSED") - await clmm_repo.close_position(position.position_address) - self._position_missing_strikes.pop(position.position_address, None) + await self.clmm_service.mark_position_closed(position_address) + self._position_missing_strikes.pop(position_address, None) else: - logger.debug(f"Position {position.position_address} miss " + logger.debug(f"Position {position_address} miss " f"{strikes}/{self.MISSING_STRIKES_TO_CLOSE} " f"(status: {status_code}), not closing yet") return # Other errors → skip update, don't close - logger.debug(f"Gateway error for position {position.position_address}: " + logger.debug(f"Gateway error for position {position_address}: " f"{result.get('error')} (status: {status_code})") return # Validate response has required fields if "address" not in result: - logger.warning(f"Invalid response for position {position.position_address}, missing 'address' field") + logger.warning(f"Invalid response for position {position_address}, missing 'address' field") return # Successful read: the position exists — reset the miss counter. - self._position_missing_strikes.pop(position.position_address, None) + self._position_missing_strikes.pop(position_address, None) except Exception as e: - logger.warning(f"Error fetching position {position.position_address} from Gateway: {e}") + logger.warning(f"Error fetching position {position_address} from Gateway: {e}") return # Extract current state @@ -870,7 +684,7 @@ async def _refresh_position_state(self, position: GatewayCLMMPosition, clmm_repo # If amounts are missing or None, skip update (don't assume zero) if base_amount_raw is None or quote_amount_raw is None: - logger.warning(f"Position {position.position_address} missing token amounts in response, skipping update") + logger.warning(f"Position {position_address} missing token amounts in response, skipping update") return base_token_amount = Decimal(str(base_amount_raw)) @@ -878,32 +692,28 @@ async def _refresh_position_state(self, position: GatewayCLMMPosition, clmm_repo # If Gateway confirms zero liquidity, position was closed externally if base_token_amount == 0 and quote_token_amount == 0: - logger.info(f"Position {position.position_address} has zero liquidity, marking as CLOSED") - await clmm_repo.close_position(position.position_address) + logger.info(f"Position {position_address} has zero liquidity, marking as CLOSED") + await self.clmm_service.mark_position_closed(position_address) return - # Update liquidity amounts, in_range status, and current price - await clmm_repo.update_position_liquidity( - position_address=position.position_address, - base_token_amount=base_token_amount, - quote_token_amount=quote_token_amount, - in_range=in_range, - current_price=current_price - ) - - # Update pending fees (always update to keep in sync with on-chain state) + # Liquidity, in_range, current price and the pending fees are one reading + # of the position and are written back together. base_fee_pending = Decimal(str(result.get("baseFeeAmount", 0))) quote_fee_pending = Decimal(str(result.get("quoteFeeAmount", 0))) - await clmm_repo.update_position_fees( - position_address=position.position_address, + await self.clmm_service.record_position_state( + position_address=position_address, + base_token_amount=base_token_amount, + quote_token_amount=quote_token_amount, + in_range=in_range, + current_price=current_price, base_fee_pending=base_fee_pending, - quote_fee_pending=quote_fee_pending + quote_fee_pending=quote_fee_pending, ) - logger.debug(f"Refreshed position {position.position_address}: price={current_price}, in_range={in_range}, " + logger.debug(f"Refreshed position {position_address}: price={current_price}, in_range={in_range}, " f"base={base_token_amount}, quote={quote_token_amount}, " f"base_fee={base_fee_pending}, quote_fee={quote_fee_pending}") except Exception as e: - logger.error(f"Error refreshing position state {position.position_address}: {e}", exc_info=True) + logger.error(f"Error refreshing position state {position_address}: {e}", exc_info=True) diff --git a/services/market_data_service.py b/services/market_data_service.py index 594723d2..0ed839a6 100644 --- a/services/market_data_service.py +++ b/services/market_data_service.py @@ -86,8 +86,12 @@ def __init__( self._ticker_max_age = ticker_max_age self._ticker_subscription_ttl = ticker_subscription_ttl - # Candle feeds management + # Candle feeds management. Creating a feed validates the pair over the network before + # the feed lands in _candle_feeds, so a per-key lock collapses concurrent first-touch + # callers into a single creation; without it the losers of the race stay running with + # no reference in _candle_feeds and no teardown path can ever stop them. self._candle_feeds: Dict[str, Any] = {} + self._candle_feed_locks: Dict[str, asyncio.Lock] = {} self._last_access_times: Dict[str, float] = {} self._feed_configs: Dict[str, Tuple[FeedType, Any]] = {} @@ -156,6 +160,7 @@ def stop(self): logger.error(f"Error stopping candle feed {feed_key}: {e}") self._candle_feeds.clear() + self._candle_feed_locks.clear() self._last_access_times.clear() self._feed_configs.clear() self._tickers.clear() @@ -487,13 +492,18 @@ async def get_candles_feed(self, config: CandlesConfig): ) if feed_key not in self._candle_feeds: - self.validate_connector(config.connector) - feed = CandlesFactory.get_candle(config) - await self._validate_pair(feed, config.connector, config.trading_pair) - feed.start() - self._candle_feeds[feed_key] = feed - self._feed_configs[feed_key] = (FeedType.CANDLES, config) - logger.info(f"Created candle feed: {feed_key}") + lock = self._candle_feed_locks.setdefault(feed_key, asyncio.Lock()) + async with lock: + # Validation is a couple of network round trips, so a concurrent caller may + # have created and registered the feed while we waited on the lock. + if feed_key not in self._candle_feeds: + self.validate_connector(config.connector) + feed = CandlesFactory.get_candle(config) + await self._validate_pair(feed, config.connector, config.trading_pair) + feed.start() + self._candle_feeds[feed_key] = feed + self._feed_configs[feed_key] = (FeedType.CANDLES, config) + logger.info(f"Created candle feed: {feed_key}") self._last_access_times[feed_key] = time.time() return self._candle_feeds[feed_key] @@ -527,6 +537,17 @@ async def get_candles_df( feed = await self.get_candles_feed(config) return feed.candles_df + def _discard_candle_feed_lock(self, feed_key: str) -> None: + """ + Drop the per-feed creation lock so the dict does not grow for the process lifetime. + + A held lock is kept: its holder is mid-creation and will register the feed, and + handing the next caller a fresh lock would reopen the very race the lock closes. + """ + lock = self._candle_feed_locks.get(feed_key) + if lock is not None and not lock.locked(): + del self._candle_feed_locks[feed_key] + def stop_candle_feed(self, config: CandlesConfig): """Stop a specific candle feed.""" feed_key = self._generate_feed_key( @@ -537,6 +558,7 @@ def stop_candle_feed(self, config: CandlesConfig): try: self._candle_feeds[feed_key].stop() del self._candle_feeds[feed_key] + self._discard_candle_feed_lock(feed_key) logger.info(f"Stopped candle feed: {feed_key}") except Exception as e: logger.error(f"Error stopping candle feed {feed_key}: {e}") @@ -1013,6 +1035,7 @@ def manually_cleanup_feed( if feed_type == FeedType.CANDLES and feed_key in self._candle_feeds: self._candle_feeds[feed_key].stop() del self._candle_feeds[feed_key] + self._discard_candle_feed_lock(feed_key) del self._last_access_times[feed_key] del self._feed_configs[feed_key] @@ -1052,6 +1075,7 @@ async def _cleanup_unused_feeds(self): if feed_type == FeedType.CANDLES and feed_key in self._candle_feeds: self._candle_feeds[feed_key].stop() del self._candle_feeds[feed_key] + self._discard_candle_feed_lock(feed_key) del self._last_access_times[feed_key] del self._feed_configs[feed_key] diff --git a/services/repository_service.py b/services/repository_service.py new file mode 100644 index 00000000..3a71a717 --- /dev/null +++ b/services/repository_service.py @@ -0,0 +1,75 @@ +"""Base for services that own a database session per operation. + +Routers used to open their own ``get_session_context`` blocks and construct +repositories inline, which put the persistence rules for a trade inside HTTP +handlers — and made every endpoint re-implement the "record it, but never fail the +caller over a write" policy with its own try/except and its own log message. + +This base holds both halves in one place. A subclass names its repository class and +expresses each operation as a function of the repository: + +- :meth:`_in_repo` for reads, where errors must reach the caller (a lookup that has + to answer 404 cannot be handed a default), +- :meth:`_in_repo_best_effort` for the after-the-fact recording of an on-chain + operation that already happened, where a database failure must be logged and + swallowed. + +Same shape as ``TradingHistoryService._run_in_repo``, which collapsed this scaffold +for the orders/trades/funding reads. +""" +import logging +from typing import Any, Awaitable, Callable, Optional, Type + +from database import AsyncDatabaseManager + +logger = logging.getLogger(__name__) + + +class RepositoryService: + """Owns the session lifecycle for one repository class.""" + + #: Repository constructed with the session for every operation. + repository_class: Optional[Type] = None + + def __init__(self, db_manager: AsyncDatabaseManager): + """ + Args: + db_manager: AsyncDatabaseManager for persistence (shared, created once at startup) + """ + self.db_manager = db_manager + + async def _in_repo(self, fn: Callable[[Any], Awaitable[Any]]) -> Any: + """Run ``fn`` against a freshly constructed repository inside a session. + + Any conversion to plain dicts must happen inside ``fn``: the session closes + when this returns, and ORM instances must not outlive it. + + Exceptions propagate — the caller decides what an unavailable database means + for its endpoint. + """ + async with self.db_manager.get_session_context() as session: + return await fn(self.repository_class(session)) + + async def _in_repo_best_effort( + self, + fn: Callable[[Any], Awaitable[Any]], + *, + error_message: str, + default: Any = None, + ) -> Any: + """Same as :meth:`_in_repo`, but a write failure never reaches the caller. + + The one expression of the policy every gateway write endpoint used to carry + its own copy of: the transaction is already on-chain, so failing the HTTP + request over the bookkeeping would report a failure that did not happen. + + Args: + fn: Async callable receiving the repository instance. + error_message: Prefix used when logging the swallowed exception. + default: Value returned instead, when ``fn`` raises. + """ + try: + return await self._in_repo(fn) + except Exception as e: + logger.error(f"{error_message}: {e}", exc_info=True) + return default diff --git a/services/unified_connector_service.py b/services/unified_connector_service.py index 7d78a506..403e7d48 100644 --- a/services/unified_connector_service.py +++ b/services/unified_connector_service.py @@ -965,9 +965,12 @@ async def _load_existing_orders( async with self.db_manager.get_session_context() as session: order_repo = OrderRepository(session) + # The connector needs the complete in-flight book: any cap here would + # leave live exchange orders invisible to sync and reconciliation. active_orders = await order_repo.get_active_orders( account_name=account_name, - connector_name=connector_name + connector_name=connector_name, + limit=None ) for order_record in active_orders: @@ -1004,19 +1007,30 @@ async def _sync_orders_to_database( orders_to_remove = [] try: - # Single session/transaction per connector: one SELECT per order and one commit on context exit. + # Single session/transaction per connector: one batched SELECT for the whole book, + # one flush for the statuses that actually moved, and one commit on context exit. async with self.db_manager.get_session_context() as session: order_repo = OrderRepository(session) - for client_order_id, order in list(connector.in_flight_orders.items()): + in_flight_orders = list(connector.in_flight_orders.items()) + db_orders = await order_repo.get_orders_by_client_ids( + [client_order_id for client_order_id, _ in in_flight_orders] + ) + db_orders_by_client_id = { + db_order.client_order_id: db_order for db_order in db_orders + } + + status_changed = False + for client_order_id, order in in_flight_orders: try: - db_order = await order_repo.get_order_by_client_id(client_order_id) + # Orders held in memory but absent from the DB are simply skipped. + db_order = db_orders_by_client_id.get(client_order_id) if db_order: new_status = self._map_order_state_to_status(order.current_state) if db_order.status != new_status: db_order.status = new_status - await session.flush() + status_changed = True if order.current_state in terminal_states: orders_to_remove.append(client_order_id) @@ -1024,6 +1038,9 @@ async def _sync_orders_to_database( except Exception as e: logger.error(f"Error syncing order {client_order_id}: {e}") + if status_changed: + await session.flush() + except Exception as e: logger.error(f"Error syncing orders for {account_name}/{connector_name}: {e}") @@ -1078,53 +1095,78 @@ async def reconcile_active_orders(self) -> Dict[str, int]: # Snapshot tracked orders (the set was loaded from the DB at init). tracked_orders = list(connector.in_flight_orders.values()) - # Single session/transaction per connector: every reconciled status update is - # flushed into one shared session and committed once on context exit. Each - # order's write runs inside its own savepoint so a SQLAlchemy error on one - # order is rolled back in isolation and does not poison the rest. - async with self.db_manager.get_session_context() as session: - order_repo = OrderRepository(session) - for order in tracked_orders: - client_order_id = order.client_order_id - note = None - try: - order_update = await connector._request_order_status(order) - new_state = order_update.new_state - except Exception as exc: - if connector._is_order_not_found_during_status_update_error(exc): - # The exchange does not know this order -> it is gone. - new_state = OrderState.CANCELED - note = "Reconciled on startup: order not found on exchange" - else: - # Transient/unknown error - do not touch the order. - logger.warning( - f"Could not verify order {client_order_id} on " - f"{account_name}/{connector_name}: {exc}" - ) - summary["unverified"] += 1 - continue - - db_status = self._map_order_state_to_status(new_state) - try: - async with session.begin_nested(): - await order_repo.update_order_status( - client_order_id=client_order_id, - status=db_status, - error_message=note, - ) - except Exception as exc: - # Savepoint rolled back: this order failed to persist but the - # session stays usable for the remaining orders. - logger.error(f"Failed to persist reconciled order {client_order_id}: {exc}") - summary["unverified"] += 1 - continue - - if new_state in terminal_states: - connector.in_flight_orders.pop(client_order_id, None) - summary["reconciled_terminal"] += 1 - else: - # Keep tracking so it stays cancelable via the trading endpoints. - summary["still_open"] += 1 + # Orders whose real state the exchange confirmed, as (client_order_id, state). + # Tracking is only updated once the transaction has committed. + resolved_orders = [] + unverified = 0 + try: + # Single session/transaction per connector: one batched SELECT for the whole + # tracked book, one flush for the statuses that actually moved, and one + # commit on context exit. + async with self.db_manager.get_session_context() as session: + order_repo = OrderRepository(session) + db_orders = await order_repo.get_orders_by_client_ids( + [order.client_order_id for order in tracked_orders] + ) + db_orders_by_client_id = { + db_order.client_order_id: db_order for db_order in db_orders + } + + status_changed = False + for order in tracked_orders: + client_order_id = order.client_order_id + note = None + try: + order_update = await connector._request_order_status(order) + new_state = order_update.new_state + except Exception as exc: + if connector._is_order_not_found_during_status_update_error(exc): + # The exchange does not know this order -> it is gone. + new_state = OrderState.CANCELED + note = "Reconciled on startup: order not found on exchange" + else: + # Transient/unknown error - do not touch the order. + logger.warning( + f"Could not verify order {client_order_id} on " + f"{account_name}/{connector_name}: {exc}" + ) + unverified += 1 + continue + + db_status = self._map_order_state_to_status(new_state) + # Orders tracked in memory but absent from the DB have no row to + # correct; they are still reconciled against the exchange. + db_order = db_orders_by_client_id.get(client_order_id) + if db_order is not None: + if db_order.status != db_status: + db_order.status = db_status + status_changed = True + if note and db_order.error_message != note: + db_order.error_message = note + status_changed = True + + resolved_orders.append((client_order_id, new_state)) + + if status_changed: + await session.flush() + except Exception as exc: + # The whole batch failed to persist: nothing is untracked and every order + # stays as it was, to be reconciled again on the next startup. + logger.error( + f"Failed to persist reconciled orders for " + f"{account_name}/{connector_name}: {exc}" + ) + summary["unverified"] += len(tracked_orders) + continue + + summary["unverified"] += unverified + for client_order_id, new_state in resolved_orders: + if new_state in terminal_states: + connector.in_flight_orders.pop(client_order_id, None) + summary["reconciled_terminal"] += 1 + else: + # Keep tracking so it stays cancelable via the trading endpoints. + summary["still_open"] += 1 logger.info( "Order reconciliation complete: " diff --git a/services/websocket_manager.py b/services/websocket_manager.py index eedb36e2..95069f60 100644 --- a/services/websocket_manager.py +++ b/services/websocket_manager.py @@ -194,25 +194,20 @@ async def _candles_push_loop(self, websocket: WebSocket, sub: Subscription): if new_candle: # New candle row appeared — send full history sub.last_sent_candle_ts = latest_ts - records = df.tail(sub.max_records).to_dict(orient="records") - await self._send_json(websocket, { - "type": "candles", - "subscription_id": sub.subscription_id, - "data": records, - "timestamp": time.time(), - }) + msg_type = "candles" + payload = df.tail(sub.max_records).to_dict(orient="records") else: # Live candle update — send only the last candle - last_record = df.iloc[-1].to_dict() - await self._send_json(websocket, { - "type": "candle_update", - "subscription_id": sub.subscription_id, - "data": last_record, - "timestamp": time.time(), - }) - except (WebSocketDisconnect, RuntimeError): - logger.info(f"WebSocket disconnected, stopping candles push [{sub.subscription_id}]") - break + msg_type = "candle_update" + payload = df.iloc[-1].to_dict() + sent = await self._send_or_stop(websocket, sub, msg_type, { + "type": msg_type, + "subscription_id": sub.subscription_id, + "data": payload, + "timestamp": time.time(), + }) + if not sent: + break except Exception as e: logger.error(f"Candles push error [{sub.subscription_id}]: {e}") except asyncio.CancelledError: @@ -237,15 +232,14 @@ async def _order_book_push_loop(self, websocket: WebSocket, sub: Subscription): bids = snapshot[0].head(sub.depth)[["price", "amount"]].values.tolist() asks = snapshot[1].head(sub.depth)[["price", "amount"]].values.tolist() - await self._send_json(websocket, { + sent = await self._send_or_stop(websocket, sub, "order_book", { "type": "order_book", "subscription_id": sub.subscription_id, "data": {"bids": bids, "asks": asks}, "timestamp": time.time(), }) - except (WebSocketDisconnect, RuntimeError): - logger.info(f"WebSocket disconnected, stopping order book push [{sub.subscription_id}]") - break + if not sent: + break except Exception as e: logger.error(f"Order book push error [{sub.subscription_id}]: {e}") except asyncio.CancelledError: @@ -293,15 +287,14 @@ async def _trades_push_loop(self, websocket: WebSocket, sub: Subscription): # Drain the buffer trades = sub.trade_buffer[:] sub.trade_buffer.clear() - await self._send_json(websocket, { + sent = await self._send_or_stop(websocket, sub, "trades", { "type": "trades", "subscription_id": sub.subscription_id, "data": trades, "timestamp": time.time(), }) - except (WebSocketDisconnect, RuntimeError): - logger.info(f"WebSocket disconnected, stopping trades push [{sub.subscription_id}]") - break + if not sent: + break except Exception as e: logger.error(f"Trades push error [{sub.subscription_id}]: {e}") except asyncio.CancelledError: @@ -313,6 +306,23 @@ async def _trades_push_loop(self, websocket: WebSocket, sub: Subscription): async def _send_json(websocket: WebSocket, data: dict): await websocket.send_json(data) + @staticmethod + async def _send_or_stop(websocket: WebSocket, sub: Subscription, msg_type: str, message: dict) -> bool: + """Push one frame; return False when the client is gone and the loop must stop. + + Only the send is guarded. A RuntimeError raised by a *fetch* is a service + fault (connector still initialising, transient upstream error), not a + dropped client, so it must stay with the caller's `except Exception` + branch, which logs it and retries on the next interval. Mirrors + services/executor_ws_manager.py. + """ + try: + await websocket.send_json(message) + return True + except (WebSocketDisconnect, RuntimeError): + logger.info(f"WebSocket disconnected, stopping {msg_type} push [{sub.subscription_id}]") + return False + @staticmethod async def _send_error(websocket: WebSocket, message: str): await websocket.send_json({"type": "error", "message": message}) diff --git a/setup.sh b/setup.sh index a220fb68..4afd8989 100755 --- a/setup.sh +++ b/setup.sh @@ -749,6 +749,89 @@ TAILSCALE_ENABLED=$TAILSCALE_ENABLED TAILSCALE_MODE=$TAILSCALE_MODE TAILSCALE_AUTH_KEY=$TAILSCALE_AUTH_KEY TAILSCALE_HOSTNAME=$TAILSCALE_HOSTNAME + +# --------------------------------------------------------------------------- +# Optional settings. Every line below is commented out and shows the default +# that config.py already applies -- config.py stays the source of truth, so +# uncomment a line only to override it. +# --------------------------------------------------------------------------- + +# Performance snapshots +# How often a live executor's performance is snapshotted, in seconds. +#PERFORMANCE_EXECUTOR_SNAPSHOT_INTERVAL=60 +# Delete performance snapshots (executor AND controller) older than this many +# days. 0 keeps everything forever -- raise it to cap database growth. +#PERFORMANCE_RETENTION_DAYS=0 + +# Backtesting +# How many backtests may run at once; further submissions queue. Runs are +# isolated in separate processes, so this can be raised up to the cores you are +# willing to give them. +#BACKTESTING_MAX_CONCURRENT=1 +# Wall-clock budget for one backtest; the worker is killed and the task fails past it. +#BACKTESTING_TIMEOUT_SECONDS=1800.0 +# How many finished backtests to retain before the oldest are reaped. +#BACKTESTING_MAX_RESULTS=100 +# Directory holding archived backtest payloads (inside the bots volume, so it +# survives redeploys). +#BACKTESTING_RESULTS_PATH=bots/data/backtests +# Directory holding downloaded candle history shared by backtest workers, so a +# sweep of configs over one market downloads that market once. +#BACKTESTING_CANDLES_CACHE_PATH=bots/data/backtests/candles +# How many downloaded candle ranges to keep; the least recently used are dropped +# past it. 0 disables the cache and makes every run download its own history. +#BACKTESTING_CANDLES_CACHE_ENTRIES=32 +# How long a downloaded candle range may be reused. A window ending near now is +# fetched with its last candle still forming, so an entry is refetched once it is +# older than this. +#BACKTESTING_CANDLES_CACHE_TTL_SECONDS=3600.0 + +# Market data feeds +# How often to run feed cleanup, in seconds. +#MARKET_DATA_CLEANUP_INTERVAL=300 +# How long to keep unused feeds alive, in seconds. +#MARKET_DATA_FEED_TIMEOUT=600 +# How long to wait for a candle feed to become ready, in seconds. +#MARKET_DATA_CANDLES_READY_TIMEOUT=30 +# WebSocket heartbeat interval, in seconds. +#MARKET_DATA_WS_HEARTBEAT_INTERVAL=30 +# Allowed range for a market-data WebSocket subscription update interval, in seconds. +#MARKET_DATA_WS_MIN_UPDATE_INTERVAL=0.25 +#MARKET_DATA_WS_MAX_UPDATE_INTERVAL=60.0 +# Allowed range, and the applied default, for a /ws/executors subscription +# update interval, in seconds. The floor is stricter than the market-data one +# because executor push loops hit the database instead of reading in-memory +# candles and order books. +#MARKET_DATA_WS_EXECUTOR_MIN_UPDATE_INTERVAL=0.5 +#MARKET_DATA_WS_EXECUTOR_MAX_UPDATE_INTERVAL=60.0 +#MARKET_DATA_WS_EXECUTOR_DEFAULT_UPDATE_INTERVAL=2.0 +# How often to refresh tickers from connected exchanges, in seconds. +#MARKET_DATA_TICKER_UPDATE_INTERVAL=30 +# Max age of cached tickers before an on-demand request refetches them, in seconds. +#MARKET_DATA_TICKER_MAX_AGE=60 +# How long a ticker-only connector stays in the background refresh cycle after +# its last request, in seconds. +#MARKET_DATA_TICKER_SUBSCRIPTION_TTL=600 +# How often to refresh Hyperliquid HIP-3 (builder-deployed perp dex) tickers, in +# seconds. Set to 0 to exclude HIP-3 markets. +#MARKET_DATA_HYPERLIQUID_HIP3_INTERVAL=120 + +# CORS (SEC-019). A wildcard origin must never be combined with credentials, so +# trusted origins are restricted to localhost by default. +# Explicit list of trusted origins, as JSON, e.g. ["https://dashboard.example.com"] +#CORS_ALLOW_ORIGINS=[] +# Regex matching trusted origins; an empty string disables regex matching. +#CORS_ALLOW_ORIGIN_REGEX=https?://(localhost|127\.0\.0\.1)(:\d+)? +# Allow credentialed (cookies/auth) cross-origin requests. +#CORS_ALLOW_CREDENTIALS=true +# HTTP methods and headers allowed for cross-origin requests. +#CORS_ALLOW_METHODS=["*"] +#CORS_ALLOW_HEADERS=["*"] + +# AWS (only needed for S3 archiving) +#AWS_API_KEY= +#AWS_SECRET_KEY= +#AWS_S3_DEFAULT_BUCKET_NAME= EOF # Appended rather than templated: these are the caller's tuning, not this diff --git a/test/test_active_orders_limit.py b/test/test_active_orders_limit.py new file mode 100644 index 00000000..1866a8c6 --- /dev/null +++ b/test/test_active_orders_limit.py @@ -0,0 +1,122 @@ +""" +Tests that `get_active_orders` bounds results explicitly and that the connector startup +path loads the whole in-flight book with no silent cap (CORR-107). + +Run with: pytest test/test_active_orders_limit.py -v +""" +import logging +from contextlib import asynccontextmanager +from types import SimpleNamespace +from unittest.mock import MagicMock + +import pytest + + +class FakeResult: + def __init__(self, rows): + self._rows = rows + + def scalars(self): + return self + + def all(self): + return list(self._rows) + + +class FakeSession: + """Answers a SELECT from memory, honouring whatever LIMIT the query carries.""" + + def __init__(self, rows): + self._rows = rows + self.limits = [] + + async def execute(self, statement): + limit_clause = getattr(statement, "_limit_clause", None) + limit = None if limit_clause is None else limit_clause.value + self.limits.append(limit) + return FakeResult(self._rows if limit is None else self._rows[:limit]) + + +def _rows(count, status="OPEN"): + return [SimpleNamespace(client_order_id=f"OID-{i}", status=status) for i in range(count)] + + +class TestGetActiveOrdersLimit: + @pytest.mark.asyncio + async def test_default_limit_still_bounds_the_query(self): + from database.repositories.order_repository import OrderRepository + + default_limit = OrderRepository.DEFAULT_ACTIVE_ORDERS_LIMIT + session = FakeSession(_rows(default_limit + 500)) + + orders = await OrderRepository(session).get_active_orders() + + assert session.limits == [default_limit] + assert len(orders) == default_limit + + @pytest.mark.asyncio + async def test_limit_none_returns_the_whole_book(self): + from database.repositories.order_repository import OrderRepository + + session = FakeSession(_rows(OrderRepository.DEFAULT_ACTIVE_ORDERS_LIMIT + 500)) + + orders = await OrderRepository(session).get_active_orders(limit=None) + + assert session.limits == [None] + assert len(orders) == OrderRepository.DEFAULT_ACTIVE_ORDERS_LIMIT + 500 + + @pytest.mark.asyncio + async def test_a_truncated_query_warns_with_the_count_and_the_limit(self, caplog): + from database.repositories.order_repository import OrderRepository + + session = FakeSession(_rows(30)) + + with caplog.at_level(logging.WARNING, logger="database.repositories.order_repository"): + await OrderRepository(session).get_active_orders(limit=10) + + warnings = [record.getMessage() for record in caplog.records + if record.levelno == logging.WARNING] + assert len(warnings) == 1 + assert "10 orders" in warnings[0] + assert "limit of 10" in warnings[0] + + @pytest.mark.asyncio + async def test_an_unclipped_query_stays_quiet(self, caplog): + from database.repositories.order_repository import OrderRepository + + session = FakeSession(_rows(3)) + + with caplog.at_level(logging.WARNING, logger="database.repositories.order_repository"): + await OrderRepository(session).get_active_orders(limit=10) + + assert [r for r in caplog.records if r.levelno == logging.WARNING] == [] + + +class TestStartupLoadsEveryActiveOrder: + @pytest.mark.asyncio + async def test_more_orders_than_the_default_limit_are_all_loaded(self): + pytest.importorskip("hummingbot") + from database.repositories.order_repository import OrderRepository + from services.unified_connector_service import UnifiedConnectorService + + total = OrderRepository.DEFAULT_ACTIVE_ORDERS_LIMIT + 250 + session = FakeSession(_rows(total)) + + @asynccontextmanager + async def get_session_context(): + yield session + + service = UnifiedConnectorService.__new__(UnifiedConnectorService) + service.db_manager = MagicMock() + service.db_manager.get_session_context = get_session_context + service._convert_db_order_to_in_flight = lambda record: SimpleNamespace( + client_order_id=record.client_order_id + ) + + connector = MagicMock() + connector.in_flight_orders = {} + + await service._load_existing_orders(connector, "master", "binance") + + assert session.limits == [None] + assert len(connector.in_flight_orders) == total diff --git a/test/test_amm_router_persistence.py b/test/test_amm_router_persistence.py new file mode 100644 index 00000000..7280a353 --- /dev/null +++ b/test/test_amm_router_persistence.py @@ -0,0 +1,596 @@ +"""The rows an AMM route writes must not change when the writer moves. + +Session ownership moved out of `routers/gateway_amm.py` into `GatewayAMMService`, and the +bot-run cleanup out of `routers/archived_bots.py` into `BotsOrchestrator` (ARCH-109), the +last two routers that still opened their own transactions. Those AMM routes are the only +writers of `gateway_amm_events` and `gateway_amm_positions` — the AMM history and the +DAMM v2 position book — so moving the writer is only safe if the same rows still land, +with the same values. + +These tests drive the real FastAPI routes with a fake repository behind the service and +pin every column each handler writes, the two paginated envelopes it hands back, and the +policies the move had to preserve: a persistence failure never fails a write that already +happened on-chain, and only a failure that actually reached the chain is recorded. +""" +from contextlib import asynccontextmanager +from decimal import Decimal +from types import SimpleNamespace +from unittest.mock import AsyncMock + +import pytest +from fastapi import FastAPI +from fastapi.testclient import TestClient + +from deps import get_accounts_service, get_bots_orchestrator, get_gateway_amm_service +from routers import archived_bots, gateway_amm +from services.gateway_amm_service import GatewayAMMService + +WALLET = "82SggYRE2Vo4jN4a2pk3aQ4SET4ctafZJGbowmCqyHx5" +POOL = "2sf5NYcY4zUPXUSmG6f66mskb24t5F8S11pC1Nz5nQT3" +POSITION = "9xQeWvG816bUx9EPjHmaT23yvVM2ZWbrrpZb9PusVFin" +SIGNATURE = "5xLmQ5s5xZ9jTqk3Y8bNvW2pR7cH4dF6gJ1kM3nP9qS8tU4vX6yZ2aB5cD7eF9gH1jK3zM5nP7qR9sT" +SOL = "So11111111111111111111111111111111111111112" +USDC = "EPjFWdd5AufqSSqeM2qN1xzybapC8G4wEGGkZwyTDt1v" +SYMBOLS = {SOL: "SOL", USDC: "USDC"} + +POOL_INFO = {"baseTokenAddress": SOL, "quoteTokenAddress": USDC, "price": 200.0} + + +# --------------------------------------------------------------------------- +# Fakes +# --------------------------------------------------------------------------- + +class RepoCalls(list): + """Every repository call a request made, in order.""" + + def payload(self, method): + """The single payload passed to ``method`` (fails if not called exactly once).""" + matches = [args for name, args in self if name == method] + assert len(matches) == 1, f"{method} called {len(matches)} times, expected 1" + return matches[0] + + def names(self): + return [name for name, _ in self] + + +def _repo_class(calls, position=None, events=(), positions=(), failing=False): + """A repository that records what it was asked to write.""" + + class _Repo: + def __init__(self, session): + if failing: + raise RuntimeError("database is down") + + async def get_position_by_address(self, address): + calls.append(("get_position_by_address", address)) + return position + + async def create_position(self, position_data): + calls.append(("create_position", position_data)) + return SimpleNamespace(id=7) + + async def add_to_position_amounts(self, **kwargs): + calls.append(("add_to_position_amounts", kwargs)) + return position + + async def subtract_from_position_amounts(self, **kwargs): + calls.append(("subtract_from_position_amounts", kwargs)) + return position + + async def close_position(self, address, **kwargs): + calls.append(("close_position", {"position_address": address, **kwargs})) + + async def create_event(self, event_data): + calls.append(("create_event", event_data)) + return SimpleNamespace(id=11) + + async def search_events(self, **kwargs): + calls.append(("search_events", kwargs)) + return list(events) + + async def search_positions(self, **kwargs): + calls.append(("search_positions", kwargs)) + return list(positions) + + @staticmethod + def event_to_dict(event): + return {"id": event.id} + + @staticmethod + def position_to_dict(position): + return {"position_address": position.position_address} + + return _Repo + + +def _db_manager(): + manager = SimpleNamespace() + + @asynccontextmanager + async def session_context(): + yield object() + + manager.get_session_context = session_context + return manager + + +def _service(calls, **repo_kwargs): + service = GatewayAMMService(db_manager=_db_manager()) + service.repository_class = _repo_class(calls, **repo_kwargs) + return service + + +def _accounts_service(**gateway_client_methods): + gateway_client = SimpleNamespace( + ping=AsyncMock(return_value=True), + parse_network_id=lambda network_id: tuple(network_id.split("-", 1)), + get_wallet_address_or_default=AsyncMock(return_value=WALLET), + resolve_token_symbol=AsyncMock(side_effect=lambda chain, network, address: SYMBOLS[address]), + **{name: AsyncMock(return_value=value) for name, value in gateway_client_methods.items()}, + ) + return SimpleNamespace(gateway_client=gateway_client) + + +def _client(accounts_service, amm_service=None): + app = FastAPI() + app.include_router(gateway_amm.router) + app.dependency_overrides[get_accounts_service] = lambda: accounts_service + app.dependency_overrides[get_gateway_amm_service] = lambda: amm_service + return TestClient(app, raise_server_exceptions=False) + + +def _stored_position(**overrides): + """A position row as the repository hands it back.""" + return SimpleNamespace(**{ + "position_address": POSITION, + "pool_address": POOL, + "wallet_address": WALLET, + "status": "OPEN", + "closed_at": None, + **overrides, + }) + + +# --------------------------------------------------------------------------- +# AMM add: a DAMM v2 position row and its ADD_LIQUIDITY event +# --------------------------------------------------------------------------- + +ADD_BODY = { + "connector": "meteora", + "network": "solana-mainnet-beta", + "pool_address": POOL, + "base_token_amount": 0.01, + "quote_token_amount": 2, +} + +ADD_RESULT = { + "signature": SIGNATURE, + "status": 1, + "data": { + "positionAddress": POSITION, + "positionRent": 0.05788, + "baseTokenAmountAdded": 0.0099, + "quoteTokenAmountAdded": 1.98, + "fee": 0.000011772, + }, +} + + +def test_add_liquidity_writes_the_same_position_row(): + calls = RepoCalls() + accounts_service = _accounts_service(amm_pool_info=POOL_INFO, amm_add_liquidity=ADD_RESULT) + client = _client(accounts_service, amm_service=_service(calls)) + + response = client.post("/gateway/amm/add-liquidity", json=ADD_BODY) + assert response.status_code == 200 + + assert calls.payload("create_position") == { + # Gateway names the position its transaction created; nothing else can. + "position_address": POSITION, + "pool_address": POOL, + "connector": "meteora", + "network": "solana-mainnet-beta", + "wallet_address": WALLET, + "base_token": "SOL", + "quote_token": "USDC", + "trading_pair": "SOL-USDC", + # The on-chain amounts, never the requested 0.01 / 2. + "initial_base_token_amount": 0.0099, + "initial_quote_token_amount": 1.98, + "base_token_amount": 0.0099, + "quote_token_amount": 1.98, + # Locked, not spent — kept as a Decimal so the close can be checked against it. + "position_rent": Decimal("0.05788"), + "entry_price": 200.0, + "current_price": 200.0, + } + + +def test_add_liquidity_writes_the_same_event_row(): + calls = RepoCalls() + accounts_service = _accounts_service(amm_pool_info=POOL_INFO, amm_add_liquidity=ADD_RESULT) + client = _client(accounts_service, amm_service=_service(calls)) + + client.post("/gateway/amm/add-liquidity", json=ADD_BODY) + + assert calls.payload("create_event") == { + "transaction_hash": SIGNATURE, + "connector": "meteora", + "network": "solana-mainnet-beta", + "wallet_address": WALLET, + "pool_address": POOL, + "position_address": POSITION, + "event_type": "ADD_LIQUIDITY", + "base_token_amount": 0.0099, + "quote_token_amount": 1.98, + "price": 200.0, + "gas_fee": 0.000011772, + "gas_token": "SOL", + "status": "CONFIRMED", + } + + +def test_adding_to_a_tracked_position_tops_it_up_and_reopens_it(): + calls = RepoCalls() + closed = _stored_position(status="CLOSED", closed_at="2026-08-01T00:00:00") + accounts_service = _accounts_service(amm_pool_info=POOL_INFO, amm_add_liquidity=ADD_RESULT) + client = _client(accounts_service, amm_service=_service(calls, position=closed)) + + response = client.post("/gateway/amm/add-liquidity", json={**ADD_BODY, "position_address": POSITION}) + assert response.status_code == 200 + + assert calls.payload("add_to_position_amounts") == { + "position_address": POSITION, + "base_delta": Decimal("0.0099"), + "quote_delta": Decimal("1.98"), + "entry_price": Decimal("200.0"), + } + # Capital went back into a position that had been emptied: it is open again. + assert (closed.status, closed.closed_at) == ("OPEN", None) + assert "create_position" not in calls.names() + + +def test_a_fungible_lp_add_records_the_event_and_no_position(): + # Raydium CPMM has no position identity — the event log is the entire record. + calls = RepoCalls() + accounts_service = _accounts_service(amm_pool_info=POOL_INFO, amm_add_liquidity=ADD_RESULT) + client = _client(accounts_service, amm_service=_service(calls)) + + response = client.post("/gateway/amm/add-liquidity", json={**ADD_BODY, "connector": "raydium"}) + assert response.status_code == 200 + + assert calls.names() == ["create_event"] + assert calls.payload("create_event")["position_address"] is None + + +def test_a_submitted_add_books_nothing_and_records_null_amounts(): + calls = RepoCalls() + accounts_service = _accounts_service( + amm_pool_info=POOL_INFO, + amm_add_liquidity={"signature": SIGNATURE, "status": 0}, + ) + client = _client(accounts_service, amm_service=_service(calls)) + + response = client.post("/gateway/amm/add-liquidity", json=ADD_BODY) + assert response.status_code == 200 + + assert calls.names() == ["create_event"] + event = calls.payload("create_event") + assert event["status"] == "SUBMITTED" + # Not yet confirmed: no amounts are invented, and no gas token without a fee. + assert (event["base_token_amount"], event["quote_token_amount"]) == (None, None) + assert (event["gas_fee"], event["gas_token"]) == (None, None) + + +def test_add_liquidity_answers_the_caller_even_when_the_write_fails(): + # The liquidity is in the pool; a bookkeeping failure must not be reported as a + # failed add. That policy now lives in one place, so this is what pins it. + accounts_service = _accounts_service(amm_pool_info=POOL_INFO, amm_add_liquidity=ADD_RESULT) + client = _client(accounts_service, amm_service=_service(RepoCalls(), failing=True)) + + response = client.post("/gateway/amm/add-liquidity", json=ADD_BODY) + + assert response.status_code == 200 + assert response.json()["signature"] == SIGNATURE + assert response.json()["status"] == "CONFIRMED" + + +def test_an_unreadable_pool_costs_the_price_not_the_write(): + calls = RepoCalls() + accounts_service = _accounts_service(amm_add_liquidity=ADD_RESULT) + accounts_service.gateway_client.amm_pool_info = AsyncMock(side_effect=RuntimeError("gateway down")) + client = _client(accounts_service, amm_service=_service(calls)) + + response = client.post("/gateway/amm/add-liquidity", json={**ADD_BODY, "connector": "raydium"}) + + assert response.status_code == 200 + assert calls.payload("create_event")["price"] is None + + +# --------------------------------------------------------------------------- +# AMM remove: unbooking the capital, and the close a 100% remove is +# --------------------------------------------------------------------------- + +REMOVE_BODY = { + "connector": "meteora", + "network": "solana-mainnet-beta", + "pool_address": POOL, + "position_address": POSITION, + "percentage_to_remove": 100, +} + +REMOVE_RESULT = { + "signature": SIGNATURE, + "status": 1, + "data": { + "baseTokenAmountRemoved": 0.0099, + "quoteTokenAmountRemoved": 1.98, + "positionRentRefunded": 0.05788, + "fee": 0.000005, + }, +} + + +def test_a_full_remove_unbooks_the_capital_and_closes_the_row(): + calls = RepoCalls() + accounts_service = _accounts_service(amm_pool_info=POOL_INFO, amm_remove_liquidity=REMOVE_RESULT) + client = _client(accounts_service, amm_service=_service(calls, position=_stored_position())) + + response = client.post("/gateway/amm/remove-liquidity", json=REMOVE_BODY) + assert response.status_code == 200 + + assert calls.payload("subtract_from_position_amounts") == { + "position_address": POSITION, + "base_delta": Decimal("0.0099"), + "quote_delta": Decimal("1.98"), + } + # Gateway closes the position account in the same transaction, which is what + # returns the rent — so the refund is recorded against the close. + assert calls.payload("close_position") == { + "position_address": POSITION, + "position_rent_refunded": Decimal("0.05788"), + } + assert calls.payload("create_event")["event_type"] == "REMOVE_LIQUIDITY" + + +def test_a_partial_remove_leaves_the_position_open(): + calls = RepoCalls() + accounts_service = _accounts_service( + amm_pool_info=POOL_INFO, + amm_remove_liquidity={**REMOVE_RESULT, "data": {"baseTokenAmountRemoved": 0.005, + "quoteTokenAmountRemoved": 1.0}}, + ) + client = _client(accounts_service, amm_service=_service(calls, position=_stored_position())) + + response = client.post("/gateway/amm/remove-liquidity", + json={**REMOVE_BODY, "percentage_to_remove": 50}) + assert response.status_code == 200 + + assert calls.payload("subtract_from_position_amounts")["base_delta"] == Decimal("0.005") + # The account stays open and refunds nothing. + assert "close_position" not in calls.names() + + +def test_a_remove_that_only_submitted_unbooks_nothing(): + calls = RepoCalls() + accounts_service = _accounts_service( + amm_pool_info=POOL_INFO, + amm_remove_liquidity={"signature": SIGNATURE, "status": 0}, + ) + client = _client(accounts_service, amm_service=_service(calls, position=_stored_position())) + + response = client.post("/gateway/amm/remove-liquidity", json=REMOVE_BODY) + + assert response.status_code == 200 + assert calls.names() == ["create_event"] + + +# --------------------------------------------------------------------------- +# Create-pool: the event files against the pool the response names +# --------------------------------------------------------------------------- + +def test_create_pool_records_its_event_against_the_pool_it_created(): + calls = RepoCalls() + accounts_service = _accounts_service(amm_create_pool={ + "signature": SIGNATURE, + "status": 1, + "poolAddress": POOL, + "price": 200.0, + "data": {"baseTokenAmountAdded": 0.0099, "quoteTokenAmountAdded": 1.98, "fee": 0.0002}, + }) + client = _client(accounts_service, amm_service=_service(calls)) + + response = client.post("/gateway/amm/create-pool", json={ + "connector": "meteora", + "network": "solana-mainnet-beta", + "base_token": "SOL", + "quote_token": "USDC", + "base_token_amount": 0.01, + "quote_token_amount": 2, + "extra_params": {"configAddress": POOL}, + }) + assert response.status_code == 200 + + assert calls.payload("create_event") == { + "transaction_hash": SIGNATURE, + "connector": "meteora", + "network": "solana-mainnet-beta", + "wallet_address": WALLET, + # Read from the response: the request names tokens, not a pool that did not exist. + "pool_address": POOL, + "position_address": None, + "event_type": "CREATE_POOL", + "base_token_amount": 0.0099, + "quote_token_amount": 1.98, + "price": 200.0, + "gas_fee": 0.0002, + "gas_token": "SOL", + "status": "CONFIRMED", + } + + +# --------------------------------------------------------------------------- +# The failed-write path keeps its single transaction-id parser +# --------------------------------------------------------------------------- + +@pytest.mark.asyncio +async def test_a_reverted_write_is_filed_with_the_gas_it_burned(): + from services.gateway_client import GatewayError + + calls = RepoCalls() + error = GatewayError(f"Transaction {SIGNATURE} landed on-chain but failed: 0x1771", status=500) + + await gateway_amm._record_failed_event( + _service(calls), error, event_type="ADD_LIQUIDITY", connector="meteora", + network="solana-mainnet-beta", wallet_address=WALLET, pool_address=POOL, + position_address=POSITION, + ) + + assert calls.payload("create_event") == { + "transaction_hash": SIGNATURE, + "connector": "meteora", + "network": "solana-mainnet-beta", + "wallet_address": WALLET, + "pool_address": POOL, + "position_address": POSITION, + "event_type": "ADD_LIQUIDITY", + "status": "FAILED", + "error_message": str(error), + } + + +@pytest.mark.asyncio +async def test_a_failure_that_never_reached_the_chain_writes_nothing(): + from services.gateway_client import GatewayError + + calls = RepoCalls() + + await gateway_amm._record_failed_event( + _service(calls), + GatewayError("Simulation failed: insufficient funds", status=400), + event_type="ADD_LIQUIDITY", connector="meteora", network="solana-mainnet-beta", + wallet_address=WALLET, pool_address=POOL, + ) + + assert calls.names() == [] + + +# --------------------------------------------------------------------------- +# The read envelopes the handlers used to hand-build +# --------------------------------------------------------------------------- + +def test_event_search_answers_the_same_envelope(): + calls = RepoCalls() + events = [SimpleNamespace(id=1), SimpleNamespace(id=2)] + client = _client(_accounts_service(), amm_service=_service(calls, events=events)) + + response = client.post("/gateway/amm/events/search?limit=2000&offset=10&event_type=CREATE_POOL") + + assert response.status_code == 200 + assert response.json() == { + "data": [{"id": 1}, {"id": 2}], + "total_count": 2, + # The query is clamped to the 1000 ceiling; the envelope echoes what was asked. + "limit": 2000, + "offset": 10, + } + assert calls.payload("search_events") == { + "connector": None, + "network": None, + "wallet_address": None, + "pool_address": None, + "event_type": "CREATE_POOL", + "status": None, + "limit": 1000, + "offset": 10, + } + + +def test_position_search_answers_the_same_envelope(): + calls = RepoCalls() + positions = [_stored_position(position_address=f"POS-{i}") for i in range(2)] + client = _client(_accounts_service(), amm_service=_service(calls, positions=positions)) + + response = client.post("/gateway/amm/positions/search?wallet_address=" + WALLET) + + assert response.status_code == 200 + assert response.json() == { + "data": [{"position_address": "POS-0"}, {"position_address": "POS-1"}], + "total_count": 2, + "limit": 50, + "offset": 0, + } + assert calls.payload("search_positions")["wallet_address"] == WALLET + + +def test_an_unreachable_database_is_a_500_not_an_empty_page(): + client = _client(_accounts_service(), amm_service=_service(RepoCalls(), failing=True)) + + response = client.post("/gateway/amm/events/search") + + assert response.status_code == 500 + assert response.json()["detail"].startswith("Error searching AMM events") + + +# --------------------------------------------------------------------------- +# Archived bots: the bot-run cleanup its delete route reports +# --------------------------------------------------------------------------- + +def _archived_bots_client(orchestrator): + app = FastAPI() + app.include_router(archived_bots.router) + app.dependency_overrides[get_bots_orchestrator] = lambda: orchestrator + return TestClient(app, raise_server_exceptions=False) + + +def test_deleting_an_archived_bot_reports_the_bot_runs_it_cleaned(monkeypatch): + monkeypatch.setattr(archived_bots.fs_util, "delete_archived_bot", lambda db_path: "hummingbot-1") + orchestrator = SimpleNamespace(delete_bot_runs_for_bot=AsyncMock(return_value=3)) + + response = _archived_bots_client(orchestrator).delete("/archived-bots/hummingbot-1/data/bot.sqlite") + + assert response.status_code == 200 + assert response.json() == { + "message": "Archived bot 'hummingbot-1' deleted successfully", + "bot_name": "hummingbot-1", + "bot_runs_deleted": 3, + } + orchestrator.delete_bot_runs_for_bot.assert_awaited_once_with("hummingbot-1") + + +@pytest.mark.asyncio +async def test_a_bot_run_cleanup_failure_never_fails_the_deletion(monkeypatch): + """The files are already gone; a database problem must not report otherwise.""" + from services import bots_orchestrator as orchestrator_module + + class _ExplodingRepo: + def __init__(self, session): + raise RuntimeError("database is down") + + monkeypatch.setattr(orchestrator_module, "BotRunRepository", _ExplodingRepo) + orchestrator = orchestrator_module.BotsOrchestrator.__new__(orchestrator_module.BotsOrchestrator) + orchestrator.db_manager = _db_manager() + + assert await orchestrator.delete_bot_runs_for_bot("hummingbot-1") == 0 + + +@pytest.mark.asyncio +async def test_the_cleanup_returns_what_the_repository_deleted(monkeypatch): + from services import bots_orchestrator as orchestrator_module + + deleted = [] + + class _StubRepo: + def __init__(self, session): + pass + + async def delete_bot_runs_by_bot_name(self, bot_name): + deleted.append(bot_name) + return 2 + + monkeypatch.setattr(orchestrator_module, "BotRunRepository", _StubRepo) + orchestrator = orchestrator_module.BotsOrchestrator.__new__(orchestrator_module.BotsOrchestrator) + orchestrator.db_manager = _db_manager() + + assert await orchestrator.delete_bot_runs_for_bot("hummingbot-1") == 2 + assert deleted == ["hummingbot-1"] diff --git a/test/test_archived_bot_zero_fill.py b/test/test_archived_bot_zero_fill.py new file mode 100644 index 00000000..f71e9d30 --- /dev/null +++ b/test/test_archived_bot_zero_fill.py @@ -0,0 +1,150 @@ +"""An archived bot that never filled a trade must not 500 its performance routes. + +A bot whose order size sits below the exchange's minimum notional has every order +rejected and archives with an empty TradeFill table. `pd.read_sql_query` gives a +zero-row table `object` columns -- there is nothing to infer a dtype from -- and +`.cumsum()` refuses object dtype however empty the frame is, so +/archived-bots/{db}/performance and /summary both answered 500 with +"cumsum is not supported for object dtype" while /executors and /orders on the same +archive answered fine. + +Run with: pytest test/test_archived_bot_zero_fill.py -v +""" +import sqlite3 + +import pytest + +pytest.importorskip("pandas") + +from utils.hummingbot_database_reader import HummingbotDatabase # noqa: E402 + +# The columns each reader touches, in the types the real hummingbot schema declares. +SCHEMA = """ +CREATE TABLE TradeFill ( + config_file_path VARCHAR, strategy VARCHAR, market VARCHAR, symbol VARCHAR, + base_asset VARCHAR, quote_asset VARCHAR, timestamp BIGINT, order_id VARCHAR, + trade_type VARCHAR, order_type VARCHAR, price BIGINT, amount BIGINT, + leverage INTEGER, trade_fee VARCHAR, trade_fee_in_quote BIGINT, + exchange_trade_id VARCHAR, position VARCHAR +); +CREATE TABLE "Order" ( + id VARCHAR, config_file_path VARCHAR, strategy VARCHAR, market VARCHAR, + symbol VARCHAR, base_asset VARCHAR, quote_asset VARCHAR, + creation_timestamp BIGINT, order_type VARCHAR, amount BIGINT, leverage INTEGER, + price BIGINT, last_status VARCHAR, last_update_timestamp BIGINT, + exchange_order_id VARCHAR, position VARCHAR +); +CREATE TABLE OrderStatus (id INTEGER, order_id VARCHAR, timestamp BIGINT, status VARCHAR); +CREATE TABLE Executors ( + id VARCHAR, timestamp BIGINT, type VARCHAR, close_timestamp BIGINT, + close_type INTEGER, status INTEGER, config VARCHAR, net_pnl_pct FLOAT, + net_pnl_quote FLOAT, cum_fees_quote FLOAT, filled_amount_quote FLOAT, + is_active BOOLEAN, is_trading BOOLEAN, custom_info VARCHAR, controller_id VARCHAR +); +CREATE TABLE Controllers (id VARCHAR, controller_id VARCHAR, timestamp BIGINT, config VARCHAR); +CREATE TABLE Position ( + id INTEGER, controller_id VARCHAR, connector_name VARCHAR, trading_pair VARCHAR, + side VARCHAR, timestamp BIGINT, volume_traded_quote BIGINT, amount BIGINT, + breakeven_price BIGINT, unrealized_pnl_quote BIGINT, cum_fees_quote BIGINT +); +""" + + +@pytest.fixture +def zero_fill_db(tmp_path): + """An archive of a bot whose every order was rejected: rows in Order, none in TradeFill.""" + path = tmp_path / "zero_fill_repro.sqlite" + connection = sqlite3.connect(path) + connection.executescript(SCHEMA) + connection.execute( + 'INSERT INTO "Order" (id, config_file_path, strategy, market, symbol, base_asset, ' + 'quote_asset, creation_timestamp, order_type, amount, leverage, price, last_status, ' + 'last_update_timestamp, exchange_order_id, position) VALUES ' + "('o-1', 'conf.yml', 'v2_with_controllers', 'binance_perpetual_testnet', 'BTC-USDT', " + "'BTC', 'USDT', 1757000000000, 'LIMIT', 235000, 5, 100000000000, 'FAILED', " + "1757000000000, NULL, 'NIL')" + ) + connection.commit() + connection.close() + return str(path) + + +class TestTheReaderSurvivesAnEmptyTable: + def test_trade_fills_of_an_empty_table_do_not_raise(self, zero_fill_db): + """The reported crash: cumsum on the object-dtype columns of a zero-row read.""" + trade_fills = HummingbotDatabase(zero_fill_db).get_trade_fills() + + assert len(trade_fills) == 0 + assert "cum_fees_in_quote" in trade_fills.columns + assert "trade_fee" in trade_fills.columns + + def test_the_scaled_columns_are_numeric_whatever_the_row_count(self, zero_fill_db): + """Object dtype is the defect, not the symptom: it breaks any later arithmetic.""" + db = HummingbotDatabase(zero_fill_db) + + trade_fills = db.get_trade_fills() + for column in ["amount", "price", "trade_fee_in_quote"]: + assert trade_fills[column].dtype != object, column + + positions = db.get_positions() + for column in ["volume_traded_quote", "amount", "breakeven_price", + "unrealized_pnl_quote", "cum_fees_quote"]: + assert positions[column].dtype != object, column + + def test_performance_of_a_zero_fill_archive_is_an_empty_frame_not_a_raise(self, zero_fill_db): + performance = HummingbotDatabase(zero_fill_db).calculate_trade_based_performance() + + assert len(performance) == 0 + + def test_a_filled_archive_still_scales_and_accumulates(self, zero_fill_db): + """The coercion must not change what a table with rows in it reports.""" + connection = sqlite3.connect(zero_fill_db) + connection.execute( + "INSERT INTO TradeFill (config_file_path, strategy, market, symbol, base_asset, " + "quote_asset, timestamp, order_id, trade_type, order_type, price, amount, " + "leverage, trade_fee, trade_fee_in_quote, exchange_trade_id, position) VALUES " + "('conf.yml', 'v2', 'binance_perpetual_testnet', 'BTC-USDT', 'BTC', 'USDT', " + "1757000000000, 'o-1', 'BUY', 'LIMIT', 100000000000, 2000000, 5, '{}', 500000, " + "'t-1', 'NIL'), " + "('conf.yml', 'v2', 'binance_perpetual_testnet', 'BTC-USDT', 'BTC', 'USDT', " + "1757000060000, 'o-2', 'SELL', 'LIMIT', 101000000000, 2000000, 5, '{}', 700000, " + "'t-2', 'NIL')" + ) + connection.commit() + connection.close() + + trade_fills = HummingbotDatabase(zero_fill_db).get_trade_fills() + + assert len(trade_fills) == 2 + assert trade_fills["price"].tolist() == [100000.0, 101000.0] + assert trade_fills["amount"].tolist() == [2.0, 2.0] + assert trade_fills["cum_fees_in_quote"].tolist() == [0.5, 1.2] + + +class TestTheRoutesAnswerTheZeroFillArchive: + """The two routes David saw 500, exercised through their own bodies.""" + + @pytest.fixture(autouse=True) + def _resolve_to_the_fixture(self, monkeypatch, zero_fill_db): + import routers.archived_bots as archived_bots + + monkeypatch.setattr(archived_bots, "_validate_db_path", lambda db_path: zero_fill_db) + + @pytest.mark.asyncio + async def test_summary_counts_the_orders_and_reports_no_trades(self): + from routers.archived_bots import get_database_summary + + summary = await get_database_summary("archived/zero_fill_repro.sqlite") + + assert summary["total_trades"] == 0 + assert summary["total_orders"] == 1 + assert summary["trading_pairs"] == ["BTC-USDT"] + + @pytest.mark.asyncio + async def test_performance_says_there_are_no_trades(self): + from routers.archived_bots import get_database_performance + + performance = await get_database_performance("archived/zero_fill_repro.sqlite") + + assert performance["performance_data"] == [] + assert performance["error"] == "No trades found in database" diff --git a/test/test_backtest_candle_cache.py b/test/test_backtest_candle_cache.py new file mode 100644 index 00000000..da488511 --- /dev/null +++ b/test/test_backtest_candle_cache.py @@ -0,0 +1,355 @@ +""" +Tests that repeated backtests over one market download its candle history once (PERF-112), +without handing two runs anything mutable to share. + +History: a backtest used to run on a service-wide BacktestingEngineBase whose data provider +cached downloaded candles across runs -- and corrupted concurrent runs doing it (CORR-060). +ARCH-063 gave every run its own engine in its own process, which fixed that structurally and +lost the cache with it: an optimizer sweeping N configs over one market downloaded that +market N times. + +The cache is back, outside the process, and holds only candle *data*: a reader gets its own +unpickled copy, so there is still no engine, provider or controller reachable from two runs. +These tests pin both halves -- that the download happens once, and that what is shared cannot +be mutated by one run into another's view -- plus the two bounds a cache has to have: it is +capped in size, and it never answers for a range it does not actually hold. + +The repo has no async test setup, so coroutines are driven with asyncio.run(). + +Run with: pytest test/test_backtest_candle_cache.py -v +""" +import asyncio +import multiprocessing +import os +import time +from pathlib import Path + +import pandas as pd +import pytest + +from services.backtesting_service import _install_candle_cache, _run_backtest_blocking +from services.candles_cache import CandlesCache, cache_key + +WINDOW = (1_700_000_000, 1_700_086_400) + + +class _CandlesConfig: + """The fields of hummingbot's CandlesConfig that decide what a download fetches.""" + + def __init__(self, connector="binance", trading_pair="BTC-USDT", interval="1m", max_records=500): + self.connector = connector + self.trading_pair = trading_pair + self.interval = interval + self.max_records = max_records + + +def _frame(seed=0.0): + return pd.DataFrame({"timestamp": [WINDOW[0], WINDOW[1]], "close": [100.0 + seed, 101.0 + seed]}) + + +class _Provider: + """Stands in for BacktestingDataProvider: what the wrapper touches, and a counted download.""" + + def __init__(self, downloads, start_time=WINDOW[0], end_time=WINDOW[1]): + self.start_time = start_time + self.end_time = end_time + self.candles_feeds = {} + self.downloads = downloads + + @staticmethod + def _generate_candle_feed_key(config): + return f"{config.connector}_{config.trading_pair}_{config.interval}" + + async def get_candles_feed(self, config): + self.downloads.append(self._generate_candle_feed_key(config)) + return _frame(len(self.downloads)) + + +class _CoveringProvider(_Provider): + """Upstream's own rule: a feed already held for this market answers, whatever was asked.""" + + async def get_candles_feed(self, config): + held = self.candles_feeds.get(self._generate_candle_feed_key(config)) + if held is not None: + return held + return await super().get_candles_feed(config) + + +def _cache(tmp_path, max_entries=32, ttl_seconds=3600.0): + return CandlesCache(path=str(tmp_path / "candles"), max_entries=max_entries, ttl_seconds=ttl_seconds) + + +def _run(cache, downloads, config=None, **provider_kwargs): + """One backtest's worth of candle fetching, on a provider that has never seen the market.""" + provider = _Provider(downloads, **provider_kwargs) + _install_candle_cache(provider, cache) + return asyncio.run(provider.get_candles_feed(config or _CandlesConfig())), provider + + +def _entries(tmp_path): + return list((tmp_path / "candles").glob("*.pkl")) + + +class TestRepeatedRunsDownloadOnce: + def test_a_sweep_over_one_market_downloads_it_once(self, tmp_path): + """The item's case: N configs, one market, one window -- one download.""" + cache, downloads = _cache(tmp_path), [] + frames = [_run(cache, downloads)[0] for _ in range(7)] + + assert len(downloads) == 1, f"a sweep of 7 runs downloaded {len(downloads)} times" + for frame in frames[1:]: + pd.testing.assert_frame_equal(frame, frames[0]) + + def test_the_provider_still_holds_the_feed_under_its_own_key(self, tmp_path): + """A cache hit must leave the provider looking exactly like a download did.""" + cache, downloads = _cache(tmp_path), [] + _run(cache, downloads) + _, provider = _run(cache, downloads) + + assert list(provider.candles_feeds) == ["binance_BTC-USDT_1m"] + pd.testing.assert_frame_equal(provider.candles_feeds["binance_BTC-USDT_1m"], _frame(1)) + + def test_a_run_asking_twice_reads_once(self, tmp_path): + """The second ask inside one run is answered from the run's own memo, not the disk.""" + cache, downloads = _cache(tmp_path), [] + provider = _Provider(downloads) + _install_candle_cache(provider, cache) + + async def scenario(): + first = await provider.get_candles_feed(_CandlesConfig()) + second = await provider.get_candles_feed(_CandlesConfig()) + return first, second + + first, second = asyncio.run(scenario()) + assert first is second + assert len(downloads) == 1 + + +class TestNothingMutableIsShared: + """The guarantee CORR-060 established and this cache must not undo.""" + + def test_each_run_gets_its_own_copy_of_the_frame(self, tmp_path): + cache, downloads = _cache(tmp_path), [] + first, _ = _run(cache, downloads) + second, _ = _run(cache, downloads) + + assert first is not second, "two runs were handed the same mutable frame" + + def test_one_run_mutating_its_frame_cannot_reach_another(self, tmp_path): + cache, downloads = _cache(tmp_path), [] + first, _ = _run(cache, downloads) + first.loc[0, "close"] = -999.0 + + second, _ = _run(cache, downloads) + assert second.loc[0, "close"] == 100.0 + 1 + assert len(downloads) == 1 + + +class TestItNeverAnswersForARangeItDoesNotHold: + @pytest.mark.parametrize("provider_kwargs", [ + {"start_time": WINDOW[0] - 86_400}, # window starts earlier + {"end_time": WINDOW[1] + 86_400}, # window ends later + ]) + def test_a_different_window_is_a_miss(self, tmp_path, provider_kwargs): + cache, downloads = _cache(tmp_path), [] + _run(cache, downloads) + _run(cache, downloads, **provider_kwargs) + + assert len(downloads) == 2, "a window the cache does not hold was served from it anyway" + + @pytest.mark.parametrize("config_kwargs", [ + {"connector": "kucoin"}, + {"trading_pair": "ETH-USDT"}, + {"interval": "5m"}, + {"max_records": 1000}, # decides how far before the window the fetch reaches back + ]) + def test_a_different_feed_is_a_miss(self, tmp_path, config_kwargs): + cache, downloads = _cache(tmp_path), [] + _run(cache, downloads) + _run(cache, downloads, config=_CandlesConfig(**config_kwargs)) + + assert len(downloads) == 2, f"{config_kwargs} was served another feed's candles" + + def test_a_feed_reused_inside_a_run_is_not_stored_as_another_key(self, tmp_path): + """Upstream reuses a feed it already holds whatever max_records asked for. + + That frame was fetched for the narrower request, so caching it under the wider one + would serve a later run less history than it asked for. Only real downloads are kept. + """ + cache, downloads = _cache(tmp_path), [] + provider = _CoveringProvider(downloads) + _install_candle_cache(provider, cache) + + async def a_run_that_asks_narrow_then_wide(): + await provider.get_candles_feed(_CandlesConfig(max_records=500)) + await provider.get_candles_feed(_CandlesConfig(max_records=1000)) + + asyncio.run(a_run_that_asks_narrow_then_wide()) + assert len(downloads) == 1 # upstream's own in-run reuse is left alone + + _run(cache, downloads, config=_CandlesConfig(max_records=1000)) + assert len(downloads) == 2, "a run was handed a frame fetched for a shorter lookback" + + def test_an_entry_past_its_freshness_bound_is_refetched(self, tmp_path): + """A window ending near now is fetched with its last candle still forming.""" + cache, downloads = _cache(tmp_path, ttl_seconds=0.0), [] + _run(cache, downloads) + time.sleep(0.01) + _run(cache, downloads) + + assert len(downloads) == 2 + + +class TestBounds: + def test_distinct_markets_cannot_grow_the_store_without_limit(self, tmp_path): + cache, downloads = _cache(tmp_path, max_entries=3), [] + for n in range(12): + _run(cache, downloads, config=_CandlesConfig(trading_pair=f"T{n}-USDT")) + + assert len(downloads) == 12 + assert len(_entries(tmp_path)) == 3, "the candle cache grew past its cap" + + def test_the_least_recently_used_entry_is_the_one_dropped(self, tmp_path): + cache, downloads = _cache(tmp_path, max_entries=2), [] + configs = [_CandlesConfig(trading_pair=f"T{n}-USDT") for n in range(3)] + _run(cache, downloads, config=configs[0]) + _run(cache, downloads, config=configs[1]) + time.sleep(0.01) + _run(cache, downloads, config=configs[0]) # touch the oldest, making it the newest + _run(cache, downloads, config=configs[2]) # evicts one + + assert len(_entries(tmp_path)) == 2 + _run(cache, downloads, config=configs[0]) + assert len(downloads) == 3, "the entry that was just used is the one that got dropped" + + def test_a_disabled_cache_downloads_every_time(self, tmp_path): + """The operator escape hatch: BACKTESTING_CANDLES_CACHE_ENTRIES=0.""" + cache, downloads = _cache(tmp_path, max_entries=0), [] + _run(cache, downloads) + _run(cache, downloads) + + assert len(downloads) == 2 + assert not (tmp_path / "candles").exists(), "a disabled cache still wrote to disk" + + +class TestTheCacheIsNeverAFailureMode: + def test_an_unreadable_entry_is_a_miss_not_a_crash(self, tmp_path): + cache, downloads = _cache(tmp_path), [] + _run(cache, downloads) + for entry in _entries(tmp_path): + entry.write_bytes(b"not a pickle") + + frame, _ = _run(cache, downloads) + assert len(downloads) == 2 + pd.testing.assert_frame_equal(frame, _frame(2)) + + def test_an_unwritable_store_still_runs_the_backtest(self, tmp_path): + cache = CandlesCache(path=str(tmp_path / "candles"), max_entries=4, ttl_seconds=3600.0) + (tmp_path / "candles").chmod(0o500) + try: + downloads = [] + frame, _ = _run(cache, downloads) + assert len(downloads) == 1 + pd.testing.assert_frame_equal(frame, _frame(1)) + finally: + (tmp_path / "candles").chmod(0o700) + + +# -- single flight across processes -- +# +# Runs are separate processes, so the workers that miss the same key at once are in separate +# processes too. Defined at module scope because a spawned child has to import them. + + +def _download_once(cache_path, key, marker_dir): + """A worker's first touch of a market: miss, take the lock, re-check, download.""" + cache = CandlesCache(path=cache_path, max_entries=8, ttl_seconds=3600.0) + if cache.get(key) is None: + with cache.single_flight(key): + if cache.get(key) is None: + Path(marker_dir, f"{os.getpid()}").write_text("") + time.sleep(0.3) # stands in for the multi-second historical download + cache.put(key, _frame()) + + +def test_workers_that_miss_the_same_market_together_download_once(tmp_path): + """Without single flight, raising BACKTESTING_MAX_CONCURRENT re-multiplies the download.""" + marker_dir = tmp_path / "downloads" + marker_dir.mkdir() + key = cache_key("binance", "BTC-USDT", "1m", 500, *WINDOW) + + ctx = multiprocessing.get_context("spawn") + workers = [ctx.Process(target=_download_once, args=(str(tmp_path / "candles"), key, str(marker_dir))) + for _ in range(4)] + for worker in workers: + worker.start() + for worker in workers: + worker.join(60) + + assert [w.exitcode for w in workers] == [0, 0, 0, 0] + downloads = list(marker_dir.iterdir()) + assert len(downloads) == 1, f"{len(downloads)} workers downloaded the same market at once" + + +# -- the real classes -- + + +class TestAgainstTheRealEngine: + def test_the_data_provider_routes_every_download_through_what_the_cache_wraps(self, tmp_path): + """Upstream's initialize_candles_feed must go through the instance attribute we wrap.""" + from hummingbot.data_feed.candles_feed.data_types import CandlesConfig + from hummingbot.strategy_v2.backtesting.backtesting_data_provider import BacktestingDataProvider + + config = CandlesConfig(connector="binance", trading_pair="BTC-USDT", interval="1m") + cache, downloads = _cache(tmp_path), [] + + def _fetch(provider): + provider.update_backtesting_time(*WINDOW) + + async def fake_download(_config): + downloads.append(_config.trading_pair) + return _frame() + + provider.get_candles_feed = fake_download # stands in for the exchange + _install_candle_cache(provider, cache) + asyncio.run(provider.initialize_candles_feed(config)) + return provider + + _fetch(BacktestingDataProvider(connectors={})) + provider = _fetch(BacktestingDataProvider(connectors={})) + + assert len(downloads) == 1, "the real provider downloaded again on a cache hit" + assert provider.get_candles_df("binance", "BTC-USDT", "1m").empty is False + + def test_the_worker_installs_the_cache_on_the_engine_it_builds(self, tmp_path, monkeypatch): + """The wiring: two worker runs of the same config download the history once.""" + import hummingbot.strategy_v2.backtesting.backtesting_engine_base as engine_module + + downloads = [] + + class _FakeEngine: + def __init__(self): + self.backtesting_data_provider = _Provider(downloads) + + @classmethod + def get_controller_config_instance_from_dict(cls, config_data, controllers_module): + return config_data + + async def run_backtesting(self, controller_config, trade_cost, start, end, backtesting_resolution): + await self.backtesting_data_provider.get_candles_feed(_CandlesConfig()) + return { + "executors": [], + "results": {"sharpe_ratio": None, "net_pnl": 0.0}, + "processed_data": {"features": _frame()}, + } + + monkeypatch.setattr(engine_module, "BacktestingEngineBase", _FakeEngine) + cache = _cache(tmp_path) + config = {"config": {"controller_name": "x"}, "start_time": WINDOW[0], "end_time": WINDOW[1]} + + first = _run_backtest_blocking(config, "conf", "controllers", cache=cache) + second = _run_backtest_blocking(config, "conf", "controllers", cache=cache) + + assert len(downloads) == 1, "the worker did not route its download through the cache" + assert first["results"]["sharpe_ratio"] == 0 # payload shaping is untouched + assert second["processed_data"] == first["processed_data"] diff --git a/test/test_backtest_concurrency.py b/test/test_backtest_concurrency.py new file mode 100644 index 00000000..837c2cc2 --- /dev/null +++ b/test/test_backtest_concurrency.py @@ -0,0 +1,172 @@ +""" +Tests that two backtests in flight at once do not corrupt each other's results, and that +no more of them run at once than the configured cap allows. + +History: the service used to keep a single BacktestingEngineBase so its data provider could +cache downloaded candles across runs, but run_backtesting builds the run on that shared +instance -- time window, controller, resolution -- and then suspends on the multi-second +candle download. A second run entering that window overwrote the first run's state and both +returned silently wrong numbers, with no exception to notice. CORR-060 closed that with a +lock around the run; ARCH-063 then moved each run into its own worker process, which owns +its engine outright, so the two runs have no reachable shared state to corrupt at all. The +lock is gone because there is nothing left for it to guard. + +These tests therefore still assert the CORR-060 guarantee -- every run's result carries its +own config and nothing else's -- but against the process mechanism, and they add the cap: +concurrency is now a deliberate, configured number rather than an accident of a lock. + +The worker below stands in for the real one so no market data is needed. It runs in a real +spawned child, so the isolation and the cap are exercised for real; it just counts how many +siblings are alive alongside it (one marker file per live worker) instead of simulating. + +The repo has no async test setup, so coroutines are driven with asyncio.run(). + +Run with: pytest test/test_backtest_concurrency.py -v +""" +import asyncio +import os +import pickle +import time +from pathlib import Path + +from services.backtesting_service import BacktestingService + +HOLD_SECONDS = 0.4 + + +def _echo_worker(config, controllers_path, controllers_module, out_path): + """Worker target: echo the run's own config back, and report the peak liveness seen.""" + live_dir = Path(config["live_dir"]) + live_dir.mkdir(parents=True, exist_ok=True) + marker = live_dir / str(os.getpid()) + marker.write_text("") + peak = 0 + try: + deadline = time.monotonic() + HOLD_SECONDS + while time.monotonic() < deadline: + peak = max(peak, len(list(live_dir.iterdir()))) + time.sleep(0.02) + finally: + marker.unlink(missing_ok=True) + + tag = config["config"]["id"] + result = { + "executors": [], + # Same shape as a real payload: {column: {row_index: value}}. + "processed_data": { + "controller_id": {0: tag}, + "resolution": {0: config["backtesting_resolution"]}, + "start": {0: config["start_time"]}, + }, + "results": {"sharpe_ratio": 0, "controller_id": tag, "peak_concurrency": peak}, + "position_holds": [], + "position_held_timeseries": [], + "pnl_timeseries": [], + } + with open(out_path, "wb") as fh: + pickle.dump({"ok": True, "result": result}, fh) + + +class StubService(BacktestingService): + """BacktestingService driving the stub worker, so no market data is needed.""" + + def __init__(self, tmp_path, max_concurrent=1): + super().__init__( + max_results=10, + results_path=str(tmp_path / "results"), + max_concurrent=max_concurrent, + timeout_seconds=30, + ) + self._worker_target = _echo_worker + + +def _config(tmp_path, tag, start, resolution): + return { + "config": {"id": tag, "controller_name": "ema_trend_v1"}, + "start_time": start, + "end_time": start + 3600, + "backtesting_resolution": resolution, + "live_dir": str(tmp_path / "live"), + } + + +def _tags(result): + """The identifying values the run carried, read back out of its own result. + + processed_data is a one-row frame rendered as {column: {row_index: value}}; the row key + is an int in memory and a string once the archive has round-tripped through JSON, so the + single value is taken positionally. + """ + features = result["processed_data"] + return ( + next(iter(features["controller_id"].values())), + next(iter(features["resolution"].values())), + next(iter(features["start"].values())), + result["results"]["controller_id"], + ) + + +def test_concurrent_executions_each_return_their_own_config(tmp_path): + async def scenario(): + service = StubService(tmp_path) + first, second = await asyncio.gather( + service._execute_backtest(_config(tmp_path, "alpha", 1000, "1m")), + service._execute_backtest(_config(tmp_path, "beta", 5000, "1h")), + ) + assert _tags(first) == ("alpha", "1m", 1000, "alpha") + assert _tags(second) == ("beta", "1h", 5000, "beta") + # No two runs were ever in flight at once under the default cap of one. + assert first["results"]["peak_concurrency"] == 1 + assert second["results"]["peak_concurrency"] == 1 + + asyncio.run(scenario()) + + +def test_sync_run_and_background_task_do_not_cross_configs(tmp_path): + """POST /backtesting/run and POST /backtesting/tasks share the one service singleton.""" + + async def scenario(): + service = StubService(tmp_path) + task = service.submit_task(_config(tmp_path, "background", 7000, "3m")) + sync_result = await service.run_backtest_sync(_config(tmp_path, "foreground", 2000, "5m")) + await task._asyncio_task + + assert _tags(sync_result) == ("foreground", "5m", 2000, "foreground") + # The archived task keeps only its metrics resident; the bulk is rehydrated from disk. + stored = service.get_task_payload(task.task_id)["result"] + assert _tags(stored) == ("background", "3m", 7000, "background") + assert sync_result["results"]["peak_concurrency"] == 1 + + asyncio.run(scenario()) + + +def test_submissions_beyond_the_cap_queue_instead_of_all_running(tmp_path): + """Four submissions against a cap of one: only one worker is ever alive.""" + + async def scenario(): + service = StubService(tmp_path, max_concurrent=1) + tasks = [service.submit_task(_config(tmp_path, f"run{i}", 1000 * i, "1m")) for i in range(4)] + await asyncio.gather(*(t._asyncio_task for t in tasks)) + for task in tasks: + result = service.get_task_payload(task.task_id)["result"] + assert result["results"]["peak_concurrency"] == 1 + assert result["results"]["controller_id"] == task.config["config"]["id"] + + asyncio.run(scenario()) + + +def test_the_cap_is_what_serializes_not_luck(tmp_path): + """With a cap of two, two runs really do overlap -- so the cap above is doing the work.""" + + async def scenario(): + service = StubService(tmp_path, max_concurrent=2) + first, second = await asyncio.gather( + service._execute_backtest(_config(tmp_path, "alpha", 1000, "1m")), + service._execute_backtest(_config(tmp_path, "beta", 5000, "1h")), + ) + assert max(first["results"]["peak_concurrency"], second["results"]["peak_concurrency"]) == 2 + # Overlapping did not let either run pick up the other's config. + assert _tags(first) == ("alpha", "1m", 1000, "alpha") + assert _tags(second) == ("beta", "1h", 5000, "beta") + + asyncio.run(scenario()) diff --git a/test/test_backtest_offloading.py b/test/test_backtest_offloading.py new file mode 100644 index 00000000..e1e7fd36 --- /dev/null +++ b/test/test_backtest_offloading.py @@ -0,0 +1,206 @@ +""" +Tests that a backtest runs off the API's event loop, can actually be stopped, and gives up +when it overruns its wall-clock budget. + +The incident these pin down: run_backtesting is awaitable, but past the candle download it +is an uninterrupted CPU loop over every candle. Awaited inline it ran on the API's only +loop thread, which sat at 100% for hours while every endpoint -- /docs included -- stopped +answering, and DELETE /backtesting/tasks/{id} reported CANCELLED while the computation kept +going, because asyncio delivers a cancellation at an await and there was none to deliver it +at. Only docker restart cleared it. + +Each test uses a real spawned worker with a cheap target, so what is exercised is the actual +process supervision -- the poll loop, the terminate, the budget -- and not a mock of it. + +The repo has no async test setup, so coroutines are driven with asyncio.run(). + +Run with: pytest test/test_backtest_offloading.py -v +""" +import asyncio +import os +import pickle +import time +from pathlib import Path + +import pytest + +from services.backtesting_service import BacktestingService, BacktestTaskStatus + +RESULT = { + "executors": [], + "processed_data": {}, + "results": {"sharpe_ratio": 0, "net_pnl": 0.0}, + "position_holds": [], + "position_held_timeseries": [], + "pnl_timeseries": [], +} + + +def _burn_worker(config, controllers_path, controllers_module, out_path): + """Worker target: hold a core the way the real simulation does, then return a payload.""" + deadline = time.monotonic() + float(config["burn_seconds"]) + spins = 0 + while time.monotonic() < deadline: + spins += 1 + with open(out_path, "wb") as fh: + pickle.dump({"ok": True, "result": {**RESULT, "results": {**RESULT["results"], "spins": spins}}}, fh) + + +def _runaway_worker(config, controllers_path, controllers_module, out_path): + """Worker target: announce its pid, then never finish -- the run that needs stopping.""" + Path(config["pid_file"]).write_text(str(os.getpid())) + deadline = time.monotonic() + 60 + while time.monotonic() < deadline: + pass + with open(out_path, "wb") as fh: + pickle.dump({"ok": True, "result": RESULT}, fh) + + +def _service(tmp_path, timeout_seconds=30): + service = BacktestingService( + max_results=10, + results_path=str(tmp_path / "results"), + max_concurrent=1, + timeout_seconds=timeout_seconds, + ) + return service + + +def _config(tmp_path, **extra): + return { + "config": {"id": "c1", "controller_name": "ema_trend_v1"}, + "start_time": 1000, + "end_time": 4600, + "backtesting_resolution": "1m", + **extra, + } + + +def _is_alive(pid): + """True while the pid exists. The parent reaps its worker, so a dead one is gone, not a zombie.""" + try: + os.kill(pid, 0) + except (ProcessLookupError, PermissionError) as e: + return isinstance(e, PermissionError) + return True + + +async def _wait_for(predicate, timeout=5.0): + deadline = time.monotonic() + timeout + while time.monotonic() < deadline: + if predicate(): + return True + await asyncio.sleep(0.02) + return False + + +def test_a_running_backtest_does_not_block_the_event_loop(tmp_path): + """The loop keeps serving while a backtest burns a core: the work is not on this thread.""" + + async def scenario(): + service = _service(tmp_path) + service._worker_target = _burn_worker + + ticks = 0 + + async def heartbeat(): + """Stands in for every other endpoint the API is supposed to keep answering.""" + nonlocal ticks + while True: + await asyncio.sleep(0.01) + ticks += 1 + + beat = asyncio.create_task(heartbeat()) + started = time.monotonic() + result = await service._execute_backtest(_config(tmp_path, burn_seconds=1.0)) + elapsed = time.monotonic() - started + beat.cancel() + + assert result["results"]["spins"] > 0 # the worker really did burn the time + assert elapsed >= 1.0 + # Inline on the loop the heartbeat would have been frozen for the whole run; the + # bar is deliberately far below the ~100 ticks a free loop manages in a second. + assert ticks >= 30, f"event loop only ticked {ticks} times during a {elapsed:.1f}s backtest" + + asyncio.run(scenario()) + + +def test_cancelling_a_task_actually_kills_the_computation(tmp_path): + """DELETE /backtesting/tasks/{id} stops the work, instead of only relabelling the task.""" + + async def scenario(): + service = _service(tmp_path) + service._worker_target = _runaway_worker + pid_file = tmp_path / "worker.pid" + + task = service.submit_task(_config(tmp_path, pid_file=str(pid_file))) + assert await _wait_for(pid_file.exists, timeout=20), "worker never started" + pid = int(pid_file.read_text()) + assert _is_alive(pid) + + assert service.cancel_task(task.task_id) is True + with pytest.raises(asyncio.CancelledError): + await task._asyncio_task + + assert task.status == BacktestTaskStatus.CANCELLED + assert not _is_alive(pid), "the worker survived the cancellation and is still burning a core" + + asyncio.run(scenario()) + + +def test_a_backtest_past_its_budget_is_terminated_and_reported(tmp_path): + """No operator intervention: the run is killed and the task explains why it failed.""" + + async def scenario(): + # The budget covers the whole run, worker start-up included, so it has to leave room + # for the child to come up -- a spawned interpreter takes the best part of a second. + service = _service(tmp_path, timeout_seconds=4.0) + service._worker_target = _runaway_worker + pid_file = tmp_path / "worker.pid" + + task = service.submit_task(_config(tmp_path, pid_file=str(pid_file))) + await task._asyncio_task + assert pid_file.exists(), "worker never started" + pid = int(pid_file.read_text()) + + assert task.status == BacktestTaskStatus.FAILED + assert "budget" in (task.error or "") + assert not _is_alive(pid), "the timed-out worker was left running" + + asyncio.run(scenario()) + + +def test_a_worker_that_dies_is_reported_as_a_failure(tmp_path): + """A worker killed from outside (OOM, operator) fails the task rather than hanging it.""" + + async def scenario(): + service = _service(tmp_path) + service._worker_target = _runaway_worker + pid_file = tmp_path / "worker.pid" + + task = service.submit_task(_config(tmp_path, pid_file=str(pid_file))) + assert await _wait_for(pid_file.exists, timeout=20), "worker never started" + os.kill(int(pid_file.read_text()), 9) + + await task._asyncio_task + assert task.status == BacktestTaskStatus.FAILED + assert "without producing a result" in (task.error or "") + + asyncio.run(scenario()) + + +def test_worker_files_do_not_accumulate(tmp_path): + """Each run cleans up after itself, and a crash's leftovers go on the next start.""" + + async def scenario(): + service = _service(tmp_path) + service._worker_target = _burn_worker + await service._execute_backtest(_config(tmp_path, burn_seconds=0.05)) + assert list(service._results_dir.glob(".worker-*.pkl")) == [] + + stale = service._results_dir / ".worker-deadbeef.pkl" + stale.write_bytes(b"leftover") + _service(tmp_path) + assert not stale.exists() + + asyncio.run(scenario()) diff --git a/test/test_backtesting_run_status_codes.py b/test/test_backtesting_run_status_codes.py new file mode 100644 index 00000000..055c8ce4 --- /dev/null +++ b/test/test_backtesting_run_status_codes.py @@ -0,0 +1,97 @@ +""" +Regression tests for POST /backtesting/run reporting failures by status code. + +The route used to wrap the whole call in `except Exception: return {"error": str(e)}`, +so every failure -- a bad controller config, a dead worker, a run abandoned at the +wall-clock budget -- came back as HTTP 200 with an error string in the body. The +api-client raises only on non-2xx, so it handed that dict straight to the caller and +nothing inspected it: the failure was silent end to end. + +Pinned here: a run that raises answers non-2xx, a run terminated by the budget answers +504 with the budget message intact (log-scraping still matches), a successful run is +still a plain 200 payload, and the /tasks submission path is untouched. + +Run with: pytest test/test_backtesting_run_status_codes.py -v +""" +from types import SimpleNamespace +from unittest.mock import AsyncMock + +import pytest +from fastapi import FastAPI +from fastapi.testclient import TestClient + +from services.backtesting_service import BacktestTaskStatus, BacktestTimeout + +TIMEOUT_MESSAGE = "Backtest exceeded its wall-clock budget of 1800s and was terminated" + +RESULT = { + "executors": [], + "processed_data": {}, + "results": {"sharpe_ratio": 0, "net_pnl": 1.5}, + "position_holds": [], + "position_held_timeseries": [], + "pnl_timeseries": [], +} + + +def _body(): + return { + "start_time": 1735689600, + "end_time": 1738368000, + "backtesting_resolution": "1m", + "trade_cost": 0.0006, + "config": {"controller_name": "pmm_simple"}, + } + + +@pytest.fixture +def client_for(): + from deps import get_backtesting_service + from routers import backtesting + + app = FastAPI() + app.include_router(backtesting.router) + + def build(service): + app.dependency_overrides[get_backtesting_service] = lambda: service + return TestClient(app, raise_server_exceptions=False) + + return build + + +def _service(**kwargs): + return SimpleNamespace(**kwargs) + + +def test_run_failure_is_not_a_200(client_for): + client = client_for(_service( + run_backtest_sync=AsyncMock(side_effect=RuntimeError("ValueError: unknown controller 'nope'")), + )) + response = client.post("/backtesting/run", json=_body()) + assert response.status_code == 500 + assert "unknown controller" in response.json()["detail"] + + +def test_run_timeout_is_504_with_the_budget_message(client_for): + client = client_for(_service( + run_backtest_sync=AsyncMock(side_effect=BacktestTimeout(TIMEOUT_MESSAGE)), + )) + response = client.post("/backtesting/run", json=_body()) + assert response.status_code == 504 + # The text is unchanged so existing log-scraping still matches. + assert response.json()["detail"] == TIMEOUT_MESSAGE + + +def test_successful_run_is_still_a_plain_200_payload(client_for): + client = client_for(_service(run_backtest_sync=AsyncMock(return_value=RESULT))) + response = client.post("/backtesting/run", json=_body()) + assert response.status_code == 200 + assert response.json()["results"]["net_pnl"] == 1.5 + + +def test_task_submission_is_unchanged(client_for): + task = SimpleNamespace(task_id="abc12345", status=BacktestTaskStatus.PENDING) + client = client_for(_service(submit_task=lambda config: task)) + response = client.post("/backtesting/tasks", json=_body()) + assert response.status_code == 200 + assert response.json() == {"task_id": "abc12345", "status": "pending"} diff --git a/test/test_bot_run_stopped_at_is_utc.py b/test/test_bot_run_stopped_at_is_utc.py new file mode 100644 index 00000000..c13c096f --- /dev/null +++ b/test/test_bot_run_stopped_at_is_utc.py @@ -0,0 +1,121 @@ +"""`stopped_at` is written as an aware UTC instant, on both paths that write it. + +`update_bot_run_stopped` used `datetime.utcnow()` -- correct in value, naive in type -- +against a `TIMESTAMP(timezone=True)` column, so the driver stored it as if it were +already in the session's local timezone and the row landed the server's UTC offset +behind the real stop time (8 hours, on the machine this was reported from). + +stop-and-archive masked it: `update_bot_run_archived` runs afterwards with an aware +value and overwrites the wrong one. A bot stopped through plain +`POST /bot-orchestration/stop-bot` and never archived kept the skew forever, and both +run duration and the attribution of performance-history windows to a run are read off +this field. + +Run with: pytest test/test_bot_run_stopped_at_is_utc.py -v +""" +import ast +from datetime import datetime, timedelta, timezone +from pathlib import Path +from types import SimpleNamespace + +import pytest + +from database.repositories.bot_run_repository import BotRunRepository + +REPOSITORIES = Path(__file__).resolve().parent.parent / "database" / "repositories" + + +class _FakeSession: + """Just the surface these two methods touch.""" + + def __init__(self, bot_run): + self.bot_run = bot_run + + async def execute(self, statement): + return SimpleNamespace(scalar_one_or_none=lambda: self.bot_run) + + async def flush(self): + pass + + async def refresh(self, instance): + pass + + +def _bot_run(): + return SimpleNamespace( + bot_name="tz_repro-20260908-233145", + run_status="RUNNING", + deployment_status="DEPLOYED", + stopped_at=None, + final_status=None, + error_message=None, + ) + + +@pytest.mark.asyncio +async def test_stopping_a_bot_records_an_aware_utc_instant(): + bot_run = _bot_run() + + await BotRunRepository(_FakeSession(bot_run)).update_bot_run_stopped("tz_repro") + + assert bot_run.run_status == "STOPPED" + assert bot_run.stopped_at.tzinfo is not None, ( + "naive datetime on a TIMESTAMP(timezone=True) column: it is stored as local time" + ) + assert bot_run.stopped_at.utcoffset() == timedelta(0) + assert abs(bot_run.stopped_at - datetime.now(timezone.utc)) < timedelta(seconds=10) + + +@pytest.mark.asyncio +async def test_an_errored_stop_records_it_the_same_way(): + """The error branch writes the same field and had the same defect.""" + bot_run = _bot_run() + + await BotRunRepository(_FakeSession(bot_run)).update_bot_run_stopped( + "tz_repro", error_message="container exited" + ) + + assert bot_run.run_status == "ERROR" + assert bot_run.stopped_at.utcoffset() == timedelta(0) + + +@pytest.mark.asyncio +async def test_archiving_agrees_with_stopping(): + """The archive path was already correct; it is what masked the bug on stop-and-archive. + + Pinned so the two never drift apart again: whichever call lands last must write the + same kind of instant. + """ + bot_run = _bot_run() + repository = BotRunRepository(_FakeSession(bot_run)) + + await repository.update_bot_run_stopped("tz_repro") + stopped_at = bot_run.stopped_at + await repository.update_bot_run_archived("tz_repro") + + assert bot_run.stopped_at.utcoffset() == stopped_at.utcoffset() == timedelta(0) + assert abs(bot_run.stopped_at - stopped_at) < timedelta(seconds=10) + + +def test_no_repository_writes_a_naive_utcnow(): + """The class of bug, not just the instance. + + Every timestamp column in database/models.py is TIMESTAMP(timezone=True), so + `datetime.utcnow()` anywhere in a repository is a row that will be read back at the + server's UTC offset. Nothing in the code or its tests said so until this test. + """ + offenders = [] + for path in sorted(REPOSITORIES.glob("*.py")): + tree = ast.parse(path.read_text()) + for node in ast.walk(tree): + if ( + isinstance(node, ast.Call) + and isinstance(node.func, ast.Attribute) + and node.func.attr == "utcnow" + ): + offenders.append(f"{path.name}:{node.lineno}") + + assert not offenders, ( + f"naive utcnow() written to timezone-aware columns at: {', '.join(offenders)} -- " + f"use datetime.now(timezone.utc)" + ) diff --git a/test/test_bot_runs_payload.py b/test/test_bot_runs_payload.py new file mode 100644 index 00000000..46688615 --- /dev/null +++ b/test/test_bot_runs_payload.py @@ -0,0 +1,128 @@ +""" +Tests for the bot-runs payload size fix. + +The ``final_status`` blob is ~99% of a bot run record's bytes, so the list +endpoint must omit it by default while the detail endpoint keeps returning it. + +Run with: pytest test/test_bot_runs_payload.py -v +""" +from contextlib import asynccontextmanager +from datetime import datetime +from types import SimpleNamespace + +import pytest + +from services.bots_orchestrator import BotsOrchestrator + + +def _make_run(run_id: int = 1): + """A BotRun-shaped stub with every field the serializer reads.""" + return SimpleNamespace( + id=run_id, + bot_name=f"bot-{run_id}", + instance_name=f"hummingbot-bot-{run_id}", + deployed_at=datetime(2026, 7, 30, 12, 0, 0), + stopped_at=None, + strategy_type="controller", + strategy_name="pmm_simple", + config_name="conf.yml", + account_name="master_account", + image_version="latest", + deployment_status="ARCHIVED", + run_status="STOPPED", + deployment_config={"script": "v2_with_controllers.py"}, + final_status={"performance": ["x" * 10000]}, + error_message=None, + ) + + +class _StubOrchestrator(BotsOrchestrator): + """Real orchestrator methods, but built without Docker/MQTT side effects.""" + + def __init__(self, runs): # noqa: D107 - deliberately skips BotsOrchestrator.__init__ + self.runs = runs + self.db_manager = SimpleNamespace(get_session_context=self._session) + + @asynccontextmanager + async def _session(self): + yield None + + +@pytest.fixture +def patched_repo(monkeypatch): + """Patch BotRunRepository so the orchestrator methods hit stub data.""" + runs = [_make_run(1), _make_run(2)] + + class _StubRepo: + def __init__(self, session): + pass + + async def get_bot_runs(self, **kwargs): + return runs + + async def get_bot_run_by_id(self, bot_run_id): + return next((r for r in runs if r.id == bot_run_id), None) + + monkeypatch.setattr("services.bots_orchestrator.BotRunRepository", _StubRepo) + return runs + + +class TestSerializer: + """Direct tests of the parameterized serializer.""" + + def test_includes_final_status_by_default(self): + serialized = BotsOrchestrator._serialize_bot_run(_make_run()) + assert "final_status" in serialized + + def test_omits_final_status_when_disabled(self): + serialized = BotsOrchestrator._serialize_bot_run(_make_run(), include_final_status=False) + assert "final_status" not in serialized + + def test_other_fields_are_unchanged_when_omitting_final_status(self): + run = _make_run() + full = BotsOrchestrator._serialize_bot_run(run) + slim = BotsOrchestrator._serialize_bot_run(run, include_final_status=False) + + assert set(full) - set(slim) == {"final_status"} + assert all(slim[key] == full[key] for key in slim) + # deployment_config stays in the slim payload on purpose (~359 B/record). + assert slim["deployment_config"] == run.deployment_config + + +class TestOrchestratorPaths: + """The list path drops the blob; the detail path keeps it.""" + + @pytest.mark.asyncio + async def test_list_omits_final_status_by_default(self, patched_repo): + orchestrator = _StubOrchestrator(patched_repo) + result = await orchestrator.get_bot_runs() + + assert len(result) == 2 + assert all("final_status" not in run for run in result) + + @pytest.mark.asyncio + async def test_list_opt_in_restores_final_status(self, patched_repo): + orchestrator = _StubOrchestrator(patched_repo) + result = await orchestrator.get_bot_runs(include_final_status=True) + + assert all("final_status" in run for run in result) + + @pytest.mark.asyncio + async def test_detail_includes_final_status(self, patched_repo): + orchestrator = _StubOrchestrator(patched_repo) + result = await orchestrator.get_bot_run_by_id(1) + + assert result is not None + assert result["final_status"] == {"performance": ["x" * 10000]} + + +class TestRouterDefaults: + """The query param must default to off — that is the whole point of the fix.""" + + def test_include_final_status_defaults_to_false(self): + import inspect + + from routers.bot_orchestration import get_bot_runs + + param = inspect.signature(get_bot_runs).parameters["include_final_status"] + assert param.default is False diff --git a/test/test_candle_feed_creation_race.py b/test/test_candle_feed_creation_race.py new file mode 100644 index 00000000..e313208f --- /dev/null +++ b/test/test_candle_feed_creation_race.py @@ -0,0 +1,174 @@ +""" +Regression tests for the check-then-create race in MarketDataService.get_candles_feed. + +Creating a candle feed validates the trading pair over the network (exchange data load plus a +REST probe) *before* the feed is registered in `_candle_feeds`. Without a per-key lock two +concurrent first-touch callers both pass the `feed_key not in self._candle_feeds` guard, both +build and start a feed, and the loser of the assignment race keeps running a live exchange +subscription that no teardown path can reach, because every teardown keys off `_candle_feeds`. + +Run with: pytest test/test_candle_feed_creation_race.py -v +""" +import asyncio + +import numpy as np +import pytest +from hummingbot.data_feed.candles_feed.data_types import CandlesConfig + +from services.market_data_service import FeedType, MarketDataService + + +class FakeCandleFeed: + """Candle feed stand-in whose validation yields to the loop, exposing the race window.""" + + def __init__(self, config: CandlesConfig): + self.config = config + self.started = False + self.stopped = False + + async def initialize_exchange_data(self): + # A real initialize does network I/O; yielding twice makes the suspension deterministic. + await asyncio.sleep(0) + await asyncio.sleep(0) + + async def fetch_candles(self, end_time=None, limit=50): + await asyncio.sleep(0) + return np.zeros((limit, 6)) + + def start(self): + self.started = True + + def stop(self): + self.stopped = True + + +class FakeCandlesFactory: + """Records every feed it builds so orphaned (unreferenced) feeds can be detected.""" + + _candles_map = {"binance": FakeCandleFeed} + + def __init__(self): + self.created = [] + + def get_candle(self, config: CandlesConfig) -> FakeCandleFeed: + feed = FakeCandleFeed(config) + self.created.append(feed) + return feed + + +@pytest.fixture +def factory(monkeypatch): + fake = FakeCandlesFactory() + monkeypatch.setattr("services.market_data_service.CandlesFactory", fake) + return fake + + +@pytest.fixture +def service(): + return MarketDataService(connector_service=None) + + +CONFIG = CandlesConfig(connector="binance", trading_pair="BTC-USDT", interval="1m") + + +async def test_concurrent_first_touch_creates_exactly_one_feed(service, factory): + """Two racing callers must collapse into a single created-and-started feed.""" + feed_a, feed_b = await asyncio.gather( + service.get_candles_feed(CONFIG), + service.get_candles_feed(CONFIG), + ) + + assert len(factory.created) == 1, "the race started more than one candle feed" + assert feed_a is feed_b + assert feed_a.started is True + assert len(service._candle_feeds) == 1 + + +async def test_no_started_feed_is_left_unreferenced(service, factory): + """Every feed that was started must be reachable through _candle_feeds, or it leaks.""" + await asyncio.gather(*(service.get_candles_feed(CONFIG) for _ in range(5))) + + referenced = set(id(feed) for feed in service._candle_feeds.values()) + orphans = [feed for feed in factory.created if feed.started and id(feed) not in referenced] + + assert orphans == [], f"{len(orphans)} started candle feed(s) left with no reference" + + +async def test_distinct_feed_keys_are_not_serialized_into_one_feed(service, factory): + """The lock is per feed_key: different pairs still get their own feeds.""" + other = CandlesConfig(connector="binance", trading_pair="ETH-USDT", interval="1m") + + feed_a, feed_b = await asyncio.gather( + service.get_candles_feed(CONFIG), + service.get_candles_feed(other), + ) + + assert feed_a is not feed_b + assert len(service._candle_feeds) == 2 + + +async def test_cached_feed_is_returned_without_recreating(service, factory): + """The second, non-concurrent call must hit the cache and pay no creation cost.""" + first = await service.get_candles_feed(CONFIG) + second = await service.get_candles_feed(CONFIG) + + assert first is second + assert len(factory.created) == 1 + + +# ==================== Lock dict pruning ==================== + +async def test_stop_candle_feed_prunes_the_lock(service, factory): + await service.get_candles_feed(CONFIG) + assert service._candle_feed_locks + + service.stop_candle_feed(CONFIG) + + assert service._candle_feed_locks == {} + + +async def test_manual_cleanup_prunes_the_lock(service, factory): + await service.get_candles_feed(CONFIG) + + service.manually_cleanup_feed(FeedType.CANDLES, "binance", "BTC-USDT", "1m") + + assert service._candle_feed_locks == {} + assert service._candle_feeds == {} + + +async def test_unused_feed_cleanup_prunes_the_lock(service, factory): + feed = await service.get_candles_feed(CONFIG) + # Age the feed past the timeout so the janitor collects it. + for key in service._last_access_times: + service._last_access_times[key] = 0.0 + + await service._cleanup_unused_feeds() + + assert feed.stopped is True + assert service._candle_feed_locks == {} + assert service._candle_feeds == {} + + +async def test_stop_service_prunes_the_locks(service, factory): + await service.get_candles_feed(CONFIG) + + service.stop() + + assert service._candle_feed_locks == {} + assert service._candle_feeds == {} + + +async def test_a_held_lock_survives_pruning(service, factory): + """ + Pruning must not hand a fresh lock to the next caller while a creator still holds the old + one: that would reopen the race the lock exists to close. + """ + feed_key = service._generate_feed_key(FeedType.CANDLES, "binance", "BTC-USDT", "1m") + lock = service._candle_feed_locks.setdefault(feed_key, asyncio.Lock()) + + async with lock: + service._discard_candle_feed_lock(feed_key) + assert service._candle_feed_locks.get(feed_key) is lock + + service._discard_candle_feed_lock(feed_key) + assert feed_key not in service._candle_feed_locks diff --git a/test/test_clmm_position_row_shape.py b/test/test_clmm_position_row_shape.py new file mode 100644 index 00000000..a5c5f25c --- /dev/null +++ b/test/test_clmm_position_row_shape.py @@ -0,0 +1,187 @@ +"""A CLMM position row must not depend on which path recorded it. + +Positions reach `gateway_clmm_positions` two ways: `/gateway/clmm/open` opens one, and +the transaction poller's discovery sweep finds one that was opened elsewhere (the UI, an +executor talking to Gateway directly). Each used to assemble the row itself, and the key +sets had drifted apart — discovery wrote `lower_bin_id`, `upper_bin_id`, `base_fee_pending` +and `quote_fee_pending`, the route wrote `position_rent`, and neither wrote the other's +columns. The same logical position therefore had two different shapes in one table, so +anything reading it back (PnL, the position search, the gas_token backfill) had to know +which writer it was looking at. + +`build_position_row` is now the only thing that decides what a new position row contains, +and these tests pin that: the two writers produce the same key set, every key is a real +column, and each writer still records what only it can know. +""" +from contextlib import asynccontextmanager +from decimal import Decimal +from types import SimpleNamespace + +import pytest + +from database.models import GatewayCLMMPosition +from services.gateway_clmm_service import GatewayCLMMService + +WALLET = "82SggYRE2Vo4jN4a2pk3aQ4SET4ctafZJGbowmCqyHx5" +POOL = "2sf5NYcY4zUPXUSmG6f66mskb24t5F8S11pC1Nz5nQT3" +POSITION = "9xQeWvG816bUx9EPjHmaT23yvVM2ZWbrrpZb9PusVFin" +SIGNATURE = "5xLmQ5s5xZ9jTqk3Y8bNvW2pR7cH4dF6gJ1kM3nP9qS8tU4vX6yZ2aB5cD7eF9gH1jK3zM5nP7qR9sT" +SOL = "So11111111111111111111111111111111111111112" +USDC = "EPjFWdd5AufqSSqeM2qN1xzybapC8G4wEGGkZwyTDt1v" +NETWORK = "solana-mainnet-beta" + + +def _service(rows): + """A GatewayCLMMService whose repository records the rows it is handed.""" + + class _Repo: + def __init__(self, session): + pass + + async def create_position(self, position_data): + rows.append(position_data) + return SimpleNamespace(id=7) + + async def create_event(self, event_data): + return SimpleNamespace(id=11) + + manager = SimpleNamespace() + + @asynccontextmanager + async def session_context(): + yield object() + + manager.get_session_context = session_context + + service = GatewayCLMMService(db_manager=manager) + service.repository_class = _Repo + return service + + +async def _route_row(): + """The row `/gateway/clmm/open` writes, via the service the route calls.""" + rows = [] + await _service(rows).record_open_position( + position_address=POSITION, + pool_address=POOL, + network=NETWORK, + connector="meteora", + wallet_address=WALLET, + trading_pair=f"{SOL}-{USDC}", + base_token=SOL, + quote_token=USDC, + lower_price=Decimal("150"), + upper_price=Decimal("250"), + entry_price=200.0, + base_amount_added=Decimal("0.0099"), + quote_amount_added=Decimal("1.98"), + position_rent=0.05788, + transaction_hash=SIGNATURE, + gas_fee=0.000011772, + gas_token="SOL", + tx_status="CONFIRMED", + ) + assert len(rows) == 1 + return rows[0] + + +# One entry of Gateway's /trading/clmm/positions-owned listing, as the sweep sees it. +DISCOVERED = { + "address": POSITION, + "poolAddress": POOL, + "baseTokenAddress": SOL, + "quoteTokenAddress": USDC, + "price": 200.0, + "lowerPrice": 150.0, + "upperPrice": 250.0, + "baseTokenAmount": 0.0099, + "quoteTokenAmount": 1.98, + "baseFeeAmount": 0.00031, + "quoteFeeAmount": 0.062, + "lowerBinId": -1841, + "upperBinId": -1789, +} + + +async def _discovered_row(pos_data=None): + """The row the poller's discovery sweep writes, via the same service.""" + rows = [] + created = await _service(rows).record_discovered_position( + pos_data=pos_data if pos_data is not None else DISCOVERED, + connector="meteora", + network=NETWORK, + wallet_address=WALLET, + ) + assert created is True + assert len(rows) == 1 + return rows[0] + + +@pytest.mark.asyncio +async def test_both_writers_produce_the_same_key_set(): + # The whole point of the item: one table, one row shape, regardless of writer. + assert set(await _route_row()) == set(await _discovered_row()) + + +@pytest.mark.asyncio +async def test_every_column_written_exists_on_the_model(): + # A unified key set is only useful if every key is real: create_position passes + # the dict straight to GatewayCLMMPosition(**row), so a stray key raises. + columns = {column.name for column in GatewayCLMMPosition.__table__.columns} + assert set(await _route_row()) <= columns + assert set(await _discovered_row()) <= columns + + +@pytest.mark.asyncio +async def test_a_discovered_position_keeps_what_only_the_chain_reports(): + # Reconciling the key sets must not cost the four columns discovery alone can + # fill: bins identify a Meteora range, and pending fees may have been accruing + # for days before the sweep first saw the position. + row = await _discovered_row() + assert row["lower_bin_id"] == -1841 + assert row["upper_bin_id"] == -1789 + assert row["base_fee_pending"] == 0.00031 + assert row["quote_fee_pending"] == 0.062 + # Unknowable for a position opened elsewhere — NULL, never a fabricated 0. + assert row["position_rent"] is None + # What the chain holds now stands in for the deposit nothing recorded. + assert row["initial_base_token_amount"] == 0.0099 + assert row["entry_price"] == 200.0 + + +@pytest.mark.asyncio +async def test_the_open_route_keeps_the_rent_and_leaves_the_rest_null(): + row = await _route_row() + assert row["position_rent"] == 0.05788 + # Gateway reports bins on a position it lists, not on the open response. + assert row["lower_bin_id"] is None + assert row["upper_bin_id"] is None + # Nothing has accrued or been collected on a position opened a moment ago. + assert row["base_fee_pending"] == 0.0 + assert row["quote_fee_pending"] == 0.0 + assert row["base_fee_collected"] == 0.0 + assert row["quote_fee_collected"] == 0.0 + # in_range stays UNKNOWN until something observes the position on-chain. + assert row["in_range"] == "UNKNOWN" + + +@pytest.mark.asyncio +async def test_both_writers_agree_on_the_derived_columns(): + # Same range, same price: the two paths reached percentage by different + # arithmetic (one on the request's Decimals, one on Gateway's floats) and must + # not disagree about the same position. + route, discovered = await _route_row(), await _discovered_row() + assert route["percentage"] == discovered["percentage"] + assert route["percentage"] == float(Decimal("100") / Decimal("150")) + assert route["lower_price"] == discovered["lower_price"] == 150.0 + assert route["upper_price"] == discovered["upper_price"] == 250.0 + assert route["status"] == discovered["status"] == "OPEN" + + +@pytest.mark.asyncio +async def test_a_discovered_position_records_whether_it_is_in_range(): + # The sweep has a live price and the range, so unlike the open route it can say. + assert (await _discovered_row())["in_range"] == "IN_RANGE" + assert (await _discovered_row({**DISCOVERED, "price": 300.0}))["in_range"] == "OUT_OF_RANGE" + # No price to compare against is not "out of range", it is not known. + assert (await _discovered_row({**DISCOVERED, "price": 0}))["in_range"] == "UNKNOWN" diff --git a/test/test_controllers_instantiate.py b/test/test_controllers_instantiate.py index b65a9543..86c847aa 100644 --- a/test/test_controllers_instantiate.py +++ b/test/test_controllers_instantiate.py @@ -85,11 +85,12 @@ def test_lp_rebalancer_reads_the_provider_it_was_given(): def test_an_untyped_provider_is_refused_rather_than_guessed(): - """Why the wheel dropped the default: Gateway 400s on a guessed trading type, so a - provider with no type has to fail at construction, not mid-operation.""" + """Gateway 400s on a guessed trading type, so a provider with no type has to fail + before the bot runs, not mid-operation. The core's parse_provider defaults an untyped + provider to "router" — wrong branch entirely for an LP controller — so the config class + owns the contract and rejects it at load, whichever wheel is installed.""" config_class = fs_util.load_controller_config_class("generic", "lp_rebalancer") fields = {**EXTRA_FIELDS["lp_rebalancer"], "lp_provider": "meteora"} - config = config_class(id="test", controller_name="lp_rebalancer", **fields) with pytest.raises(ValueError, match="expected 'name/type'"): - config.get_controller_class()(config, MagicMock(), MagicMock()) + config_class(id="test", controller_name="lp_rebalancer", **fields) diff --git a/test/test_deploy_controllers_path_traversal.py b/test/test_deploy_controllers_path_traversal.py new file mode 100644 index 00000000..ce584443 --- /dev/null +++ b/test/test_deploy_controllers_path_traversal.py @@ -0,0 +1,90 @@ +"""SEC-058: the controllers list re-read from the deployed script config is attacker-controlled. + +`DockerService.create_hummingbot_instance` reads `controllers_config` back off the script config +YAML (which any authenticated caller can write through `POST /scripts/configs/{name}`) and used to +join every entry into a source and a destination path without validating it. These tests pin that +traversal entries are skipped and that only safe single-component names are ever copied. +""" +import os + +import pytest +import yaml + +from config import settings +from models import V2ControllerDeployment +from services.docker_service import DockerService + + +def _build_bots_tree(root, controllers_config): + """Lay out the minimal bots/ tree create_hummingbot_instance expects under `root`.""" + credentials_dir = root / "bots" / "credentials" / "master_account" + credentials_dir.mkdir(parents=True) + (credentials_dir / "conf_client.yml").write_text("instance_id: placeholder\n") + + scripts_dir = root / "bots" / "conf" / "scripts" + scripts_dir.mkdir(parents=True) + (scripts_dir / "evil.yml").write_text(yaml.dump({"controllers_config": controllers_config})) + + controllers_dir = root / "bots" / "conf" / "controllers" + controllers_dir.mkdir(parents=True) + (controllers_dir / "good.yml").write_text("controller_name: good\n") + + # A file outside bots/ that a traversal entry would exfiltrate into the instance dir. + (root / "secret.txt").write_text("DB_PASSWORD=hunter2\n") + + +def _deploy(root, monkeypatch, controllers_config): + _build_bots_tree(root, controllers_config) + monkeypatch.chdir(root) + # No config password => the deployment bails out right after the config copy, so the test + # never needs a live Docker daemon. + monkeypatch.setattr(settings.security, "config_password", "") + + service = DockerService.__new__(DockerService) + service.SOURCE_PATH = str(root) + deployment = V2ControllerDeployment( + instance_name="testbot", + credentials_profile="master_account", + controllers_config=["good.yml"], + script_config="evil.yml", + ) + service.create_hummingbot_instance(deployment) + return root / "bots" / "instances" / "testbot" / "conf" / "controllers" + + +@pytest.mark.parametrize( + "entry", + [ + "../../../secret.txt", # lands on bots/instances/secret.txt with the old code + "../../../../etc/hosts", # escapes the repo entirely + "/etc/hosts", # absolute path + "subdir/good.yml", # separator inside the name + "..", + ], +) +def test_a_traversal_controller_entry_is_never_copied(tmp_path, monkeypatch, caplog, entry): + with caplog.at_level("WARNING"): + destination = _deploy(tmp_path, monkeypatch, [entry]) + + copied = sorted(p.name for p in destination.iterdir()) if destination.exists() else [] + assert copied == [] + # Nothing was written anywhere outside the instance's own conf/controllers directory. + assert not (tmp_path / "bots" / "instances" / "secret.txt").exists() + assert (tmp_path / "secret.txt").read_text() == "DB_PASSWORD=hunter2\n" + assert any(entry in record.message for record in caplog.records if record.levelname == "WARNING") + + +def test_a_non_string_controller_entry_is_skipped(tmp_path, monkeypatch): + destination = _deploy(tmp_path, monkeypatch, [{"controller": "good.yml"}]) + + copied = sorted(p.name for p in destination.iterdir()) if destination.exists() else [] + assert copied == [] + + +def test_well_formed_controller_entries_are_still_copied(tmp_path, monkeypatch): + destination = _deploy(tmp_path, monkeypatch, ["good.yml", "../../../secret.txt"]) + + assert sorted(p.name for p in destination.iterdir()) == ["good.yml"] + assert (destination / "good.yml").read_text() == "controller_name: good\n" + assert not (tmp_path / "bots" / "instances" / "secret.txt").exists() + assert os.path.exists(tmp_path / "secret.txt") diff --git a/test/test_docker_routes_report_failure.py b/test/test_docker_routes_report_failure.py new file mode 100644 index 00000000..aaa558ac --- /dev/null +++ b/test/test_docker_routes_report_failure.py @@ -0,0 +1,203 @@ +"""A Docker operation that failed must not answer 200. + +`POST /docker/stop-container/{name}` used to return the service's error verbatim, so a +container that does not exist answered **HTTP 200** with a raw docker-py string as the +body: + + "404 Client Error for http+docker://localhost/v1.55/containers/x/json: Not Found + (\"No such container: x\")" + +Two defects in one line. A caller that checks the status code -- which is how a caller +checks -- read a container that was never stopped as one that had been. And the body +published the daemon's socket URL and negotiated API version, which is a detail of how +this API talks to Docker, not an answer to anything the caller asked. + +Every container route in this router shared the shape: `except DockerException: return +str(e)`. + +Run with: pytest test/test_docker_routes_report_failure.py -v +""" +import pytest +from docker.errors import APIError, DockerException, NotFound +from fastapi import HTTPException + +from routers.docker import ( + active_containers, + available_images, + clean_exited_containers, + exited_containers, + start_container, + stop_container, +) +from services.docker_service import DockerService + +# What docker-py actually raises, socket URL and API version included. +NOT_FOUND = NotFound( + '404 Client Error for http+docker://localhost/v1.55/containers/ghost/json: ' + 'Not Found ("No such container: ghost")' +) +DAEMON_ERROR = APIError( + "500 Server Error for http+docker://localhost/v1.55/containers/ghost/stop: " + "Server Error (\"cannot stop container\")" +) + + +class _Containers: + def __init__(self, error=None, container=None): + self.error = error + self.container = container + self.pruned = 0 + + def get(self, name): + if isinstance(self.error, DockerException): + raise self.error + return self.container + + def list(self, **kwargs): + if isinstance(self.error, DockerException): + raise self.error + return [] + + def prune(self): + if isinstance(self.error, DockerException): + raise self.error + self.pruned += 1 + + +class _Container: + def __init__(self, error=None): + self.error = error + self.stopped = 0 + self.started = 0 + + def stop(self): + if self.error: + raise self.error + self.stopped += 1 + + def start(self): + if self.error: + raise self.error + self.started += 1 + + +def _service(error=None, container=None): + service = DockerService.__new__(DockerService) + service.client = type("_Client", (), {})() + service.client.containers = _Containers(error=error, container=container) + service.client.images = _Containers(error=error) + return service + + +class TestAMissingContainerIs404: + @pytest.mark.asyncio + async def test_stopping_one_raises_404(self): + with pytest.raises(HTTPException) as raised: + await stop_container("ghost", _service(error=NOT_FOUND)) + + assert raised.value.status_code == 404 + assert raised.value.detail == "No such container: ghost" + + @pytest.mark.asyncio + async def test_starting_one_raises_404(self): + with pytest.raises(HTTPException) as raised: + await start_container("ghost", _service(error=NOT_FOUND)) + + assert raised.value.status_code == 404 + + @pytest.mark.asyncio + async def test_the_detail_never_carries_the_daemon_socket_or_api_version(self): + """The reported body leaked both. Neither is an answer to what was asked.""" + with pytest.raises(HTTPException) as raised: + await stop_container("ghost", _service(error=NOT_FOUND)) + + assert "http+docker" not in raised.value.detail + assert "v1.55" not in raised.value.detail + assert "Client Error" not in raised.value.detail + + +class TestADaemonFailureIs502: + @pytest.mark.asyncio + async def test_a_stop_the_daemon_refuses_raises_502(self): + service = _service(container=_Container(error=DAEMON_ERROR)) + + with pytest.raises(HTTPException) as raised: + await stop_container("bot-1", service) + + assert raised.value.status_code == 502 + assert "http+docker" not in raised.value.detail + + @pytest.mark.asyncio + async def test_listing_containers_with_no_daemon_raises_502(self): + """It used to answer 200 with an error string where a list was documented.""" + for route in (active_containers, exited_containers): + with pytest.raises(HTTPException) as raised: + await route(None, _service(error=DAEMON_ERROR)) + assert raised.value.status_code == 502 + + @pytest.mark.asyncio + async def test_pruning_with_no_daemon_raises_502(self): + with pytest.raises(HTTPException) as raised: + await clean_exited_containers(_service(error=DAEMON_ERROR)) + + assert raised.value.status_code == 502 + + @pytest.mark.asyncio + async def test_listing_images_with_no_daemon_raises_502(self): + with pytest.raises(HTTPException) as raised: + await available_images(None, _service(error=DAEMON_ERROR)) + + assert raised.value.status_code == 502 + + +class TestSuccessStillAnswersNormally: + @pytest.mark.asyncio + async def test_a_stop_that_works_reports_it(self): + """It used to return null, which is indistinguishable from a failure body.""" + container = _Container() + service = _service(container=container) + + response = await stop_container("bot-1", service) + + assert response["success"] is True + assert container.stopped == 1 + + @pytest.mark.asyncio + async def test_a_start_that_works_reports_it(self): + container = _Container() + service = _service(container=container) + + response = await start_container("bot-1", service) + + assert response["success"] is True + assert container.started == 1 + + @pytest.mark.asyncio + async def test_listing_containers_passes_the_list_through(self): + assert await active_containers(None, _service()) == [] + + @pytest.mark.asyncio + async def test_pruning_that_works_reports_it(self): + service = _service() + + response = await clean_exited_containers(service) + + assert response["success"] is True + assert service.client.containers.pruned == 1 + + +class TestTheServiceKeepsItsInBandContract: + """stop-and-archive calls stop_container in a retry loop and reads the container's + status to decide whether it worked, so a failure here must stay a return value.""" + + def test_a_failed_stop_returns_rather_than_raises(self): + result = _service(error=NOT_FOUND).stop_container("ghost") + + assert result["success"] is False + assert result["error"] == "not_found" + + def test_a_failed_removal_returns_the_success_flag_it_always_did(self): + result = _service(error=NOT_FOUND).remove_container("ghost") + + assert result["success"] is False + assert "http+docker" not in result["message"] diff --git a/test/test_env_template_matches_config.py b/test/test_env_template_matches_config.py new file mode 100644 index 00000000..881a08c1 --- /dev/null +++ b/test/test_env_template_matches_config.py @@ -0,0 +1,139 @@ +""" +Pins the .env template written by setup.sh against config.py (READ-117). + +setup.sh's heredoc is the only template this repo has, and the generated .env is the only +place an operator ever sees a setting. Nothing derives the heredoc from config.py, so the two +drifted: five whole settings groups had no line in the template at all. These tests make that +drift fail loudly instead of silently: + +- every prefixed settings group in config.py is represented in the template, +- every setting the template documents actually exists in config.py, and +- every default the template shows is the default config.py applies. + +Run with: pytest test/test_env_template_matches_config.py -v +""" +import json +import re +from pathlib import Path +from typing import Dict, Tuple + +import pytest +from pydantic_settings import BaseSettings + +import config + +REPO_ROOT = Path(__file__).resolve().parent.parent +SETUP_SH = REPO_ROOT / "setup.sh" + +# Groups added by READ-117: documented, but every line commented out, so a freshly generated +# .env keeps producing exactly the settings config.py already applies. +COMMENTED_ONLY_PREFIXES = ("PERFORMANCE_", "BACKTESTING_", "MARKET_DATA_", "CORS_", "AWS_") + +# A variable assignment in the template: "NAME=value" (active) or "#NAME=value" (documented). +# Prose comments start with "# " and are therefore not matched. +_ASSIGNMENT = re.compile(r"^(#?)([A-Z][A-Z0-9_]*)=(.*)$") + + +def _env_template() -> str: + """The body of setup.sh's `cat > .env << EOF` heredoc.""" + source = SETUP_SH.read_text() + start = source.index("cat > .env << EOF\n") + len("cat > .env << EOF\n") + end = source.index("\nEOF\n", start) + return source[start:end] + + +def _assignments() -> Tuple[Dict[str, str], Dict[str, str]]: + """Return (active, documented) variable assignments found in the template.""" + active, documented = {}, {} + for line in _env_template().splitlines(): + match = _ASSIGNMENT.match(line) + if match is None: + continue + commented, name, value = match.groups() + (documented if commented else active)[name] = value + return active, documented + + +def _settings_groups() -> Dict[str, type]: + """Every BaseSettings group in config.py that has a non-empty env_prefix, keyed by prefix.""" + groups = {} + for attribute in vars(config).values(): + if not (isinstance(attribute, type) and issubclass(attribute, BaseSettings)): + continue + prefix = attribute.model_config.get("env_prefix", "") + if prefix: + groups[prefix] = attribute + return groups + + +def _defaults_by_env_name() -> Dict[str, str]: + """Map each prefixed setting's env var name to the default config.py applies, rendered as .env text.""" + defaults = {} + for prefix, group in _settings_groups().items(): + for field_name, field in group.model_fields.items(): + defaults[f"{prefix}{field_name.upper()}"] = _render(field.default) + return defaults + + +def _render(default) -> str: + """Render a field default the way it would be written in .env.""" + if isinstance(default, bool): + return "true" if default else "false" + if isinstance(default, (list, dict)): + return json.dumps(default) + return str(default) + + +class TestEnvTemplateMatchesConfig: + def test_every_settings_group_is_represented(self): + """A new group in config.py must show up in the template — that is exactly what drifted.""" + active, documented = _assignments() + names = set(active) | set(documented) + for prefix in _settings_groups(): + assert any(name.startswith(prefix) for name in names), ( + f"config.py defines a settings group with env_prefix {prefix!r} that setup.sh's " + f".env template never mentions; operators cannot discover it" + ) + + def test_documented_settings_exist_in_config(self): + """The template must not document a setting config.py does not have (renamed or removed).""" + _, documented = _assignments() + known = _defaults_by_env_name() + for name in documented: + assert name in known, f"setup.sh documents {name}, which no config.py settings group defines" + + def test_documented_defaults_match_config(self): + """A default shown in the template must be the default config.py applies.""" + _, documented = _assignments() + known = _defaults_by_env_name() + for name, value in documented.items(): + assert value == known[name], ( + f"setup.sh shows {name}={value} but config.py defaults to {known[name]}" + ) + + @pytest.mark.parametrize("prefix", COMMENTED_ONLY_PREFIXES) + def test_optional_groups_are_documented_but_not_set(self, prefix): + """These groups are documented only: a fresh .env must not pin a value config.py owns.""" + active, documented = _assignments() + assert any(name.startswith(prefix) for name in documented), f"no {prefix} setting is documented" + assert not [name for name in active if name.startswith(prefix)], ( + f"{prefix} settings must stay commented out so a fresh .env changes no runtime behavior" + ) + + def test_performance_retention_explains_that_zero_keeps_everything(self): + """PERFORMANCE_RETENTION_DAYS=0 grows the database forever; the template must say so.""" + template = _env_template() + assert "#PERFORMANCE_EXECUTOR_SNAPSHOT_INTERVAL=60" in template + assert "#PERFORMANCE_RETENTION_DAYS=0" in template + retention_comment = template[: template.index("#PERFORMANCE_RETENTION_DAYS=")] + assert "0 keeps everything forever" in retention_comment + + +class TestReadmePointsAtConfig: + def test_configuration_section_names_config_py_as_authoritative(self): + """The README excerpt is curated, so it has to say where the full list lives.""" + readme = (REPO_ROOT / "README.md").read_text() + section = readme[readme.index("## Configuration"):] + section = section[: section.index("\n## ")] + assert "`config.py`" in section + assert "PERFORMANCE_RETENTION_DAYS" in section diff --git a/test/test_executor_completion_before_creation_row.py b/test/test_executor_completion_before_creation_row.py new file mode 100644 index 00000000..e7eade78 --- /dev/null +++ b/test/test_executor_completion_before_creation_row.py @@ -0,0 +1,258 @@ +"""An executor's final state is never dropped because its creation row was not there yet. + +`create_executor` registers and starts the executor synchronously, then awaits +`_persist_executor_created`. The 1s control loop runs during that await, and an executor +that failed in milliseconds is already `is_closed` by then -- so the completion write can +reach the database BEFORE the creation INSERT commits, or after it failed outright. + +`ExecutorRepository.update_executor` is select-then-update and silently does nothing when +the row is missing, while the caller logged success regardless. The observable symptom was +an executor that closed instantly but sat at `status=RUNNING` / `is_active=true` in +`POST /executors/search` and in the performance report forever, with no close_type and no +PnL, until an API restart ran `cleanup_orphaned_executors`. + +What is pinned here is that the completion write is self-healing: it inserts the row from +the metadata it already carries, and a creation INSERT that lands afterwards does not +resurrect the executor to RUNNING. +""" + +from contextlib import asynccontextmanager +from decimal import Decimal +from types import SimpleNamespace +from unittest.mock import AsyncMock, MagicMock + +import pytest +from hummingbot.strategy_v2.models.executors import CloseType +from sqlalchemy import create_engine +from sqlalchemy.orm import Session +from sqlalchemy.pool import StaticPool + +from database.models import ExecutorPerformanceSnapshot, ExecutorRecord +from database.repositories.executor_repository import ExecutorRepository +from services.executor_service import ExecutorService + + +class _AsyncSessionAdapter: + """The async surface ExecutorRepository uses, over a real synchronous Session. + + aiosqlite is not installed in this environment, and the repository only ever awaits + execute/flush/refresh/begin_nested -- so this adapter runs the real SQL against the + real schema, including the UNIQUE index on executor_id and the SAVEPOINTs the repair + insert and the terminal performance snapshot rely on. Mocking the session away would + prove nothing here. + """ + + def __init__(self, session: Session): + self._session = session + + def add(self, obj): + self._session.add(obj) + + def add_all(self, objs): + self._session.add_all(objs) + + async def execute(self, statement): + return self._session.execute(statement) + + async def flush(self): + self._session.flush() + + async def refresh(self, obj): + self._session.refresh(obj) + + async def commit(self): + self._session.commit() + + async def rollback(self): + self._session.rollback() + + async def close(self): + self._session.close() + + def begin_nested(self): + nested = self._session.begin_nested() + + @asynccontextmanager + async def _savepoint(): + try: + yield nested + if nested.is_active: + nested.commit() + except Exception: + if nested.is_active: + nested.rollback() + raise + + return _savepoint() + + +@pytest.fixture +def db(): + """An in-memory database with the tables the completion path writes, plus a session factory.""" + engine = create_engine("sqlite://", connect_args={"check_same_thread": False}, + poolclass=StaticPool) + ExecutorRecord.__table__.create(engine) + # Completion also writes the executor's terminal performance snapshot, in this same + # transaction (FEAT-001). + ExecutorPerformanceSnapshot.__table__.create(engine) + + @asynccontextmanager + async def session_context(): + """Mirrors AsyncDatabaseManager.get_session_context: commit, else rollback.""" + session = Session(engine) + adapter = _AsyncSessionAdapter(session) + try: + yield adapter + await adapter.commit() + except Exception: + await adapter.rollback() + raise + finally: + await adapter.close() + + def rows(): + with Session(engine) as session: + return session.query(ExecutorRecord).all() + + try: + yield SimpleNamespace(session_context=session_context, rows=rows) + finally: + engine.dispose() + + +IDENTITY = dict( + executor_type="position_executor", + account_name="master", + connector_name="binance_perpetual", + trading_pair="BTC-USDT", + controller_id="main", +) + +FINAL_STATE = dict( + status="TERMINATED", + close_type="STOP_LOSS", + net_pnl_quote=Decimal("-12.5"), + filled_amount_quote=Decimal("400"), +) + + +# -------------------------------------------------------------------------------------- +# The repository: the completion write repairs a missing row instead of no-op'ing +# -------------------------------------------------------------------------------------- + +@pytest.mark.asyncio +async def test_updating_a_row_that_is_not_there_yet_still_does_nothing(db): + """The sharp edge the upsert exists to cover -- documented, not fixed in place.""" + async with db.session_context() as session: + assert await ExecutorRepository(session).update_executor( + executor_id="e-1", **FINAL_STATE) is None + + assert db.rows() == [] + + +@pytest.mark.asyncio +async def test_a_completion_with_no_creation_row_inserts_the_closed_record(db): + async with db.session_context() as session: + record, repaired = await ExecutorRepository(session).upsert_executor_completion( + executor_id="e-1", **IDENTITY, **FINAL_STATE) + assert repaired, "the caller must be told the row had to be repaired" + assert record is not None + + row, = db.rows() + assert row.status == "TERMINATED" + assert row.close_type == "STOP_LOSS" + assert row.net_pnl_quote == Decimal("-12.5") + assert row.closed_at is not None + assert row.trading_pair == "BTC-USDT" + + +@pytest.mark.asyncio +async def test_a_completion_with_a_creation_row_updates_it_in_place(db): + """The normal path: one row, updated, and no repair reported.""" + async with db.session_context() as session: + await ExecutorRepository(session).create_executor(executor_id="e-1", **IDENTITY) + + async with db.session_context() as session: + _, repaired = await ExecutorRepository(session).upsert_executor_completion( + executor_id="e-1", **IDENTITY, **FINAL_STATE) + assert not repaired + + row, = db.rows() + assert row.status == "TERMINATED" + assert row.close_type == "STOP_LOSS" + + +# -------------------------------------------------------------------------------------- +# The service: the phantom RUNNING executor, end to end +# -------------------------------------------------------------------------------------- + +def _service(db, metadata): + service = ExecutorService.__new__(ExecutorService) + service.db_manager = MagicMock(get_session_context=db.session_context) + service._executor_metadata = {"e-1": dict(metadata)} + service._active_executors = {} + service._lp_position_addresses = {} + service._lp_rent_recorded = set() + service._lp_rent_retry_after = {} + service._log_capture = MagicMock() + service._log_capture.get_error_count.return_value = 0 + service._record_executor_swap = AsyncMock() + return service + + +def _closed_executor(): + """An executor that failed the moment it started -- insufficient balance, say.""" + executor = MagicMock() + executor.is_closed = True + executor.status = SimpleNamespace(name="TERMINATED") + executor.close_type = CloseType.STOP_LOSS + executor.executor_info = SimpleNamespace( + net_pnl_quote=Decimal("-12.5"), + net_pnl_pct=Decimal("-0.03"), + cum_fees_quote=Decimal("0.4"), + filled_amount_quote=Decimal("400"), + ) + executor.get_custom_info.return_value = {} + return executor + + +@pytest.mark.asyncio +async def test_completion_landing_before_the_creation_insert_leaves_a_closed_record(db): + """The race itself: the control loop completes the executor mid-_persist_created.""" + metadata = {**IDENTITY, "config": {"id": "e-1"}} + service = _service(db, metadata) + executor = _closed_executor() + service._active_executors["e-1"] = executor + + # The control loop wins: it sees is_closed and persists completion first. + await service._handle_executor_completion("e-1") + + # ...and the creation INSERT that was already in flight lands afterwards. It still + # holds its metadata, which _handle_executor_completion has since dropped. + service._executor_metadata["e-1"] = dict(metadata) + await service._persist_executor_created("e-1", executor) + + row, = db.rows() + assert row.status == "TERMINATED", "the late creation insert resurrected a phantom RUNNING executor" + assert row.close_type == "STOP_LOSS" + assert row.net_pnl_quote == Decimal("-12.5") + assert row.filled_amount_quote == Decimal("400") + assert row.final_state is not None + + +@pytest.mark.asyncio +async def test_a_creation_insert_that_never_landed_is_repaired_by_the_completion(db): + """The other instance of the same root cause: _persist_executor_created swallows + every exception, so a DB hiccup used to lose the executor from the record entirely -- + and cleanup_orphaned_executors could not repair what had no row at all.""" + service = _service(db, {**IDENTITY, "config": {"id": "e-1"}}) + service._active_executors["e-1"] = _closed_executor() + + # No _persist_executor_created call at all: it failed and was swallowed. + await service._handle_executor_completion("e-1") + + row, = db.rows() + assert row.status == "TERMINATED" + assert row.close_type == "STOP_LOSS" + assert row.account_name == "master" + assert row.executor_type == "position_executor" diff --git a/test/test_executor_performance_snapshots.py b/test/test_executor_performance_snapshots.py new file mode 100644 index 00000000..2887db7e --- /dev/null +++ b/test/test_executor_performance_snapshots.py @@ -0,0 +1,1141 @@ +"""An executor's performance is a series, and one route serves it beside the controllers'. + +Before this, an executor's database row was written exactly twice -- zeros at creation, +the real figures at completion -- so a RUNNING executor's row said its PnL was zero for +its entire life, and the live numbers existed only in ExecutorService memory. Two things +followed: + + 1. There was no series to chart. A running executor could not draw its own curve, and a + closed one had a single point. + 2. **An API restart destroyed the accounting of every executor that was live.** + cleanup_orphaned_executors flipped each RUNNING row to TERMINATED/SYSTEM_CLEANUP and + touched none of the PnL columns, so the creation-time zeros became the permanent + record and /executors/performance summed them forever. + +What is pinned here: the periodic snapshot, the terminal row that ends a closed +executor's series inside the snapshot table (no join to `executors`, and no exposure to +the two paths that leave that row at zero), the reap adopting the last snapshot, the +grain-parameterized sampler, retention, and the normalized row shape both subjects map +into. +""" +from datetime import datetime, timedelta, timezone +from decimal import Decimal +from types import SimpleNamespace +from unittest.mock import MagicMock + +import pytest + +pytest.importorskip("hummingbot") + +from database.models import ControllerPerformanceSnapshot, ExecutorPerformanceSnapshot # noqa: E402 +from database.repositories.controller_performance_repository import ControllerPerformanceRepository # noqa: E402 +from database.repositories.executor_performance_repository import ExecutorPerformanceRepository # noqa: E402 +from models.performance import controller_row_to_performance_row, executor_row_to_performance_row # noqa: E402 +from services.executor_service import ExecutorService # noqa: E402 + +NOW = datetime(2026, 9, 1, 12, 0, tzinfo=timezone.utc) + + +# -------------------------------------------------------------------------------------- +# Fakes: just enough of the async session surface these repositories touch +# -------------------------------------------------------------------------------------- + +class _RecordingSession: + """Collects add_all rows; execute() returns whatever the test queued.""" + + def __init__(self, results=None): + self.added = [] + self.flushed = 0 + self._results = list(results or []) + self.executed = [] + self.savepoints = 0 + + def begin_nested(self): + session = self + + class _Savepoint: + async def __aenter__(self): + session.savepoints += 1 + return self + + async def __aexit__(self, *exc): + return False + + return _Savepoint() + + def add_all(self, rows): + self.added.extend(rows) + + def add(self, row): + self.added.append(row) + + async def execute(self, statement): + self.executed.append(statement) + if self._results: + return self._results.pop(0) + return SimpleNamespace(scalars=lambda: SimpleNamespace(all=lambda: []), rowcount=0) + + async def flush(self): + self.flushed += 1 + + +def _scalars(rows): + return SimpleNamespace(scalars=lambda: SimpleNamespace(all=lambda: rows)) + + +def _snapshot(executor_id="e-1", timestamp=NOW, is_terminal=False, close_type=None, + net_pnl=0, net_pnl_pct=0, fees=0, filled=0, status="RUNNING"): + return ExecutorPerformanceSnapshot( + timestamp=timestamp, + executor_id=executor_id, + executor_type="position_executor", + account_name="master_account", + connector_name="binance_perpetual", + trading_pair="BTC-USDT", + controller_id="main", + status=status, + close_type=close_type, + is_terminal=is_terminal, + net_pnl_quote=Decimal(str(net_pnl)), + net_pnl_pct=Decimal(str(net_pnl_pct)), + cum_fees_quote=Decimal(str(fees)), + filled_amount_quote=Decimal(str(filled)), + ) + + +def _orphan(executor_id="e-1"): + """An ExecutorRecord-shaped stub with every column the reap reads or writes.""" + return SimpleNamespace( + executor_id=executor_id, status="RUNNING", close_type=None, closed_at=None, + executor_type="position_executor", account_name="master_account", + connector_name="binance_perpetual", trading_pair="BTC-USDT", controller_id="main", + net_pnl_quote=Decimal("0"), net_pnl_pct=Decimal("0"), + cum_fees_quote=Decimal("0"), filled_amount_quote=Decimal("0"), + ) + + +def _service(active=None, **kwargs): + """An ExecutorService with nothing running but the pieces the snapshot path reads.""" + service = ExecutorService.__new__(ExecutorService) + service.db_manager = kwargs.pop("db_manager", MagicMock()) + service.default_account = "master_account" + service.performance_snapshot_interval = kwargs.pop("performance_snapshot_interval", 60.0) + service.performance_retention_days = kwargs.pop("performance_retention_days", 0) + service._active_executors = active or {} + service._executor_metadata = kwargs.pop("metadata", {}) + return service + + +def _live_executor(net_pnl="12.5", net_pnl_pct="0.0125", fees="1.8", filled="4200", + status="RUNNING"): + executor = MagicMock() + executor.executor_info = SimpleNamespace( + net_pnl_quote=Decimal(net_pnl), + net_pnl_pct=Decimal(net_pnl_pct), + cum_fees_quote=Decimal(fees), + filled_amount_quote=Decimal(filled), + ) + executor.status.name = status + return executor + + +# -------------------------------------------------------------------------------------- +# The table: narrow, typed, and with no second volume column +# -------------------------------------------------------------------------------------- + +class TestTheTable: + def test_it_has_no_separate_volume_column(self): + """filled_amount_quote IS the volume traded, on every executor type including LP. + + The split existed once (volume_traded_quote) and was removed upstream. Any design + that re-introduces it here breaks like-for-like summing on purpose. + """ + assert not hasattr(ExecutorPerformanceSnapshot, "volume_traded_quote"), ( + "the snapshot table grew a second volume column; filled_amount_quote is it" + ) + + def test_it_carries_no_custom_info_blob(self): + """These payloads carry fill_events and grid levels; two code paths strip them. + + A per-minute row is the last place to put them back. + """ + columns = {c.name for c in ExecutorPerformanceSnapshot.__table__.columns} + assert "custom_info" not in columns + assert "performance" not in columns + + def test_one_executors_series_is_indexed(self): + """WHERE executor_id = ? ORDER BY timestamp DESC is the hot query on both readers.""" + composite = { + tuple(c.name for c in index.columns) + for index in ExecutorPerformanceSnapshot.__table__.indexes + } + assert ("executor_id", "timestamp") in composite + + +# -------------------------------------------------------------------------------------- +# The dump: one row per live executor, and a database failure never breaks the loop +# -------------------------------------------------------------------------------------- + +class TestTheDump: + @pytest.mark.asyncio + async def test_it_writes_one_row_per_live_executor(self): + session = _RecordingSession() + service = _service( + active={"e-1": _live_executor(), "e-2": _live_executor(net_pnl="-3")}, + metadata={ + "e-1": _metadata("e-1"), + "e-2": _metadata("e-2"), + }, + ) + _with_session(service, session) + + await service._dump_executor_performance() + + assert len(session.added) == 2 + assert {row.executor_id for row in session.added} == {"e-1", "e-2"} + assert all(row.is_terminal is False for row in session.added) + # One timestamp for the whole batch: the points of one tick line up on a chart. + assert len({row.timestamp for row in session.added}) == 1 + + @pytest.mark.asyncio + async def test_the_row_carries_the_metrics_the_executor_reports_right_now(self): + session = _RecordingSession() + service = _service(active={"e-1": _live_executor()}, metadata={"e-1": _metadata("e-1")}) + _with_session(service, session) + + await service._dump_executor_performance() + + row = session.added[0] + assert row.net_pnl_quote == Decimal("12.5") + assert row.cum_fees_quote == Decimal("1.8") + assert row.filled_amount_quote == Decimal("4200") + assert row.status == "RUNNING" + + @pytest.mark.asyncio + async def test_an_executor_whose_info_cannot_be_read_is_skipped_not_zeroed(self): + """A fabricated zero mid-series is worse than a gap: a reader cannot tell it from + an executor that genuinely made nothing.""" + broken = MagicMock() + type(broken).executor_info = property(lambda _: (_ for _ in ()).throw(RuntimeError("boom"))) + session = _RecordingSession() + service = _service( + active={"ok": _live_executor(), "broken": broken}, + metadata={"ok": _metadata("ok"), "broken": _metadata("broken")}, + ) + _with_session(service, session) + + await service._dump_executor_performance() + + assert [row.executor_id for row in session.added] == ["ok"] + + @pytest.mark.asyncio + async def test_a_database_failure_is_logged_and_does_not_propagate(self): + """This shares the control loop's tick; a dropped snapshot must not stop executors.""" + service = _service(active={"e-1": _live_executor()}, metadata={"e-1": _metadata("e-1")}) + + class _Exploding: + def get_session_context(self): + raise RuntimeError("database is on fire") + + service.db_manager = _Exploding() + + await service._dump_executor_performance() # must not raise + + @pytest.mark.asyncio + async def test_nothing_live_means_no_session_is_opened(self): + """The tick costs nothing when there is nothing to sample.""" + db_manager = MagicMock() + service = _service(active={}, db_manager=db_manager) + + await service._dump_executor_performance() + + db_manager.get_session_context.assert_not_called() + + +# -------------------------------------------------------------------------------------- +# Retention +# -------------------------------------------------------------------------------------- + +class TestRetention: + @pytest.mark.asyncio + async def test_the_default_deletes_nothing(self): + """PERFORMANCE_RETENTION_DAYS=0 is what every existing deployment does today; an + upgrade must not start deleting an operator's history.""" + db_manager = MagicMock() + service = _service(performance_retention_days=0, db_manager=db_manager) + + await service._prune_performance_snapshots() + + db_manager.get_session_context.assert_not_called() + + @pytest.mark.asyncio + async def test_it_deletes_from_both_snapshot_tables(self): + """Retention is one policy, not two: the operator says how much performance + history to keep and both series obey.""" + session = _RecordingSession(results=[ + SimpleNamespace(rowcount=7), + SimpleNamespace(rowcount=3), + ]) + repo = ExecutorPerformanceRepository(session) + + executor_rows, controller_rows = await repo.prune_older_than(NOW - timedelta(days=30)) + + assert (executor_rows, controller_rows) == (7, 3) + targeted = [statement.table.name for statement in session.executed] + assert targeted == [ + ExecutorPerformanceSnapshot.__tablename__, + ControllerPerformanceSnapshot.__tablename__, + ] + + +# -------------------------------------------------------------------------------------- +# The sampler: the 5-minute assumption is gone +# -------------------------------------------------------------------------------------- + +class TestTheSampler: + def _series(self, count, step_minutes=1): + """`count` rows, newest first, `step_minutes` apart -- the stored order.""" + return [ + {"timestamp": (NOW - timedelta(minutes=i * step_minutes)).isoformat()} + for i in range(count) + ] + + def test_asking_for_the_native_grain_returns_every_row(self): + history = self._series(10) + sampled = ExecutorPerformanceRepository._sample_by_interval(history, 1, grain_minutes=1.0) + assert sampled == history + + def test_five_minutes_over_a_sixty_second_grain_returns_every_fifth_row(self): + """The controller repository hard-codes `interval_minutes <= 5` and `// 5`. At a + 60-second grain that returned every row untouched.""" + history = self._series(21) + sampled = ExecutorPerformanceRepository._sample_by_interval(history, 5, grain_minutes=1.0) + + assert [history.index(row) for row in sampled] == [0, 5, 10, 15, 20] + + def test_asking_for_less_than_the_grain_still_returns_the_grain(self): + """`interval` is a floor, not a guarantee.""" + history = self._series(6, step_minutes=5) + sampled = ExecutorPerformanceRepository._sample_by_interval(history, 1, grain_minutes=5.0) + assert sampled == history + + def test_a_zero_interval_cannot_reach_the_divisor(self): + repo = ExecutorPerformanceRepository(_RecordingSession(), grain_minutes=0) + assert repo.grain_minutes > 0 + + def test_one_minute_is_a_known_interval(self): + """The executor grain is finer than anything the controller map could express.""" + assert ExecutorPerformanceRepository._interval_to_minutes("1m") == 1 + + def _fleet(self, executor_ids, count, step_minutes=1): + """`count` rows per executor, interleaved and newest first -- the stored order.""" + rows = [ + {"executor_id": executor_id, + "timestamp": (NOW - timedelta(minutes=i * step_minutes)).isoformat()} + for i in range(count) + for executor_id in executor_ids + ] + return sorted(rows, key=lambda row: row["timestamp"], reverse=True) + + def test_thinning_a_fleet_keeps_every_executor(self): + """A global cursor made `interval` a rate limit on the merged series: at 5m over a + 60s grain, three executors reporting together came back as one.""" + history = self._fleet(["exec-A", "exec-B", "exec-C"], count=10) + + sampled = ExecutorPerformanceRepository._sample_by_interval(history, 5, grain_minutes=1.0) + + assert {row["executor_id"] for row in sampled} == {"exec-A", "exec-B", "exec-C"} + + def test_a_coarse_interval_still_keeps_every_executor(self): + """The wider the window, the more a global cursor swallowed. 1h over 60s kept one.""" + history = self._fleet(["exec-A", "exec-B", "exec-C"], count=10) + + sampled = ExecutorPerformanceRepository._sample_by_interval(history, 60, grain_minutes=1.0) + + assert {row["executor_id"] for row in sampled} == {"exec-A", "exec-B", "exec-C"} + # One row each: ten minutes of history cannot fill a second hourly slot. + assert len(sampled) == 3 + + def test_each_executors_own_series_is_thinned_to_the_interval(self): + """Grouping must not disable the thinning -- every scope obeys the interval.""" + history = self._fleet(["exec-A", "exec-B"], count=21) + + sampled = ExecutorPerformanceRepository._sample_by_interval(history, 5, grain_minutes=1.0) + + for executor_id in ("exec-A", "exec-B"): + times = [ + datetime.fromisoformat(row["timestamp"]) + for row in sampled if row["executor_id"] == executor_id + ] + assert len(times) == 5 + gaps = [(a - b).total_seconds() / 60 for a, b in zip(times, times[1:])] + assert all(gap >= 5 for gap in gaps) + + def test_thinning_a_fleet_stays_newest_first(self): + """Consumers page on the last row's timestamp, so the merge order still holds.""" + history = self._fleet(["exec-A", "exec-B", "exec-C"], count=12) + + sampled = ExecutorPerformanceRepository._sample_by_interval(history, 5, grain_minutes=1.0) + + times = [row["timestamp"] for row in sampled] + assert times == sorted(times, reverse=True) + + def test_the_controller_sampler_keeps_every_controller_too(self): + """The same defect, same shape, on the series the wire-compatible route serves.""" + history = sorted( + [ + {"bot_name": "bot-1", "controller_id": controller_id, + "timestamp": (NOW - timedelta(minutes=i * 5)).isoformat()} + for i in range(12) + for controller_id in (f"ctrl-{n}" for n in range(12)) + ], + key=lambda row: row["timestamp"], reverse=True, + ) + + sampled = ControllerPerformanceRepository._sample_by_interval(history, 60) + + assert len({row["controller_id"] for row in sampled}) == 12 + + +# -------------------------------------------------------------------------------------- +# Pagination: the same contract as the controller route +# -------------------------------------------------------------------------------------- + +class TestPagination: + @pytest.mark.asyncio + async def test_a_full_page_reports_more_and_cursors_on_the_last_timestamp(self): + rows = [_snapshot(timestamp=NOW - timedelta(minutes=i)) for i in range(4)] + repo = ExecutorPerformanceRepository(_RecordingSession(results=[_scalars(rows)])) + + page, next_cursor, has_more = await repo.get_performance_history(limit=3, interval="1m") + + assert len(page) == 3 + assert has_more is True + assert next_cursor == page[-1]["timestamp"] + + @pytest.mark.asyncio + async def test_the_last_page_reports_no_more_and_no_cursor(self): + rows = [_snapshot(timestamp=NOW - timedelta(minutes=i)) for i in range(2)] + repo = ExecutorPerformanceRepository(_RecordingSession(results=[_scalars(rows)])) + + page, next_cursor, has_more = await repo.get_performance_history(limit=3, interval="1m") + + assert len(page) == 2 + assert has_more is False + assert next_cursor is None + + @pytest.mark.asyncio + async def test_walking_the_cursor_returns_every_row_exactly_once(self): + """Two pages of a four-row series, driven by next_cursor as a client would.""" + all_rows = [_snapshot(timestamp=NOW - timedelta(minutes=i)) for i in range(4)] + + # Page one over-fetches by one; page two is filtered by the cursor. + session = _RecordingSession(results=[_scalars(all_rows[:3]), _scalars(all_rows[2:])]) + repo = ExecutorPerformanceRepository(session) + + first, cursor, has_more = await repo.get_performance_history(limit=2, interval="1m") + assert has_more is True + second, _, _ = await repo.get_performance_history(limit=2, cursor=cursor, interval="1m") + + seen = [row["timestamp"] for row in first + second] + assert len(seen) == len(set(seen)), "a row was served on both pages" + + +# -------------------------------------------------------------------------------------- +# The terminal row and the reap +# -------------------------------------------------------------------------------------- + +class TestTheTerminalRow: + def test_it_carries_the_close_type_and_the_final_metrics(self): + service = _service(metadata={"e-1": _metadata("e-1")}) + executor = _live_executor(status="TERMINATED") + + row = service._build_snapshot_row( + "e-1", executor, is_terminal=True, + metrics={ + "net_pnl_quote": Decimal("42"), + "net_pnl_pct": Decimal("0.05"), + "cum_fees_quote": Decimal("2"), + "filled_amount_quote": Decimal("900"), + }, + status="TERMINATED", + close_type="TAKE_PROFIT", + ) + + assert row["is_terminal"] is True + assert row["close_type"] == "TAKE_PROFIT" + assert row["net_pnl_quote"] == Decimal("42") + assert row["status"] == "TERMINATED" + + def test_a_periodic_row_never_carries_a_close_type(self): + service = _service(metadata={"e-1": _metadata("e-1")}) + + row = service._build_snapshot_row("e-1", _live_executor(), is_terminal=False) + + assert row["is_terminal"] is False + assert row["close_type"] is None + + def test_the_row_denormalizes_the_identity_so_a_series_needs_no_join(self): + service = _service(metadata={"e-1": _metadata("e-1")}) + + row = service._build_snapshot_row("e-1", _live_executor(), is_terminal=False) + + assert row["executor_type"] == "position_executor" + assert row["account_name"] == "master_account" + assert row["connector_name"] == "binance_perpetual" + assert row["trading_pair"] == "BTC-USDT" + assert row["controller_id"] == "main" + + +class TestTheReapAdoptsTheLastSnapshot: + @pytest.mark.asyncio + async def test_a_terminated_row_takes_the_last_snapshots_figures(self): + """The restart bug. Without this the row keeps its creation-time zeros forever.""" + from database.repositories.executor_repository import ExecutorRepository + + orphan = _orphan() + latest = _snapshot(net_pnl="17.25", net_pnl_pct="0.03", fees="0.9", filled="3100") + session = _RecordingSession(results=[_scalars([orphan]), _scalars([latest])]) + + cleaned = await ExecutorRepository(session).cleanup_orphaned_executors(active_executor_ids=[]) + + assert cleaned == 1 + assert orphan.status == "TERMINATED" + assert orphan.close_type == "SYSTEM_CLEANUP" + assert orphan.net_pnl_quote == Decimal("17.25") + assert orphan.cum_fees_quote == Decimal("0.9") + assert orphan.filled_amount_quote == Decimal("3100") + + @pytest.mark.asyncio + async def test_the_reap_writes_the_terminal_row_the_completion_path_would_have(self): + """The reap was the one way to reach TERMINATED with no terminal row behind it. + + /performance/latest then served the last RUNNING snapshot forever -- is_terminal + false, close_type null -- while /executors/{id} said TERMINATED/SYSTEM_CLEANUP. + """ + from database.repositories.executor_repository import ExecutorRepository + + orphan = _orphan() + latest = _snapshot(net_pnl="17.25", net_pnl_pct="0.03", fees="0.9", filled="3100") + session = _RecordingSession(results=[_scalars([orphan]), _scalars([latest])]) + + await ExecutorRepository(session).cleanup_orphaned_executors(active_executor_ids=[]) + + rows = [row for row in session.added if isinstance(row, ExecutorPerformanceSnapshot)] + assert len(rows) == 1 + row = rows[0] + assert row.is_terminal is True + assert row.status == "TERMINATED" + assert row.close_type == "SYSTEM_CLEANUP" + # It says what the record says, so the two surfaces cannot disagree. + assert row.net_pnl_quote == orphan.net_pnl_quote == Decimal("17.25") + assert row.cum_fees_quote == orphan.cum_fees_quote == Decimal("0.9") + assert row.filled_amount_quote == orphan.filled_amount_quote == Decimal("3100") + assert row.timestamp == orphan.closed_at + + @pytest.mark.asyncio + async def test_the_terminal_row_takes_its_identity_from_the_record(self): + """An executor orphaned before its first snapshot has no snapshot to read it from.""" + from database.repositories.executor_repository import ExecutorRepository + + orphan = _orphan() + session = _RecordingSession(results=[_scalars([orphan]), _scalars([])]) + + await ExecutorRepository(session).cleanup_orphaned_executors(active_executor_ids=[]) + + row = [r for r in session.added if isinstance(r, ExecutorPerformanceSnapshot)][0] + assert row.executor_id == "e-1" + assert row.executor_type == "position_executor" + assert row.account_name == "master_account" + assert row.connector_name == "binance_perpetual" + assert row.trading_pair == "BTC-USDT" + assert row.controller_id == "main" + assert row.is_terminal is True + + @pytest.mark.asyncio + async def test_a_failing_snapshot_insert_does_not_roll_back_the_reap(self): + """The reap is the accounting; the terminal row is a point on a chart.""" + from database.repositories.executor_repository import ExecutorRepository + + orphan = _orphan() + session = _RecordingSession(results=[_scalars([orphan]), _scalars([])]) + + def _explode(rows): + raise RuntimeError("snapshot insert failed") + + session.add_all = _explode + + cleaned = await ExecutorRepository(session).cleanup_orphaned_executors(active_executor_ids=[]) + + assert cleaned == 1 + assert orphan.status == "TERMINATED" + assert session.savepoints == 1 + + @pytest.mark.asyncio + async def test_an_executor_with_no_snapshot_keeps_the_old_behaviour(self): + """Created and orphaned inside one interval: there is nothing better to write.""" + from database.repositories.executor_repository import ExecutorRepository + + orphan = _orphan() + session = _RecordingSession(results=[_scalars([orphan]), _scalars([])]) + + cleaned = await ExecutorRepository(session).cleanup_orphaned_executors(active_executor_ids=[]) + + assert cleaned == 1 + assert orphan.status == "TERMINATED" + assert orphan.net_pnl_quote == Decimal("0") + # Still gets a terminal row: the zeros are what the record itself books, and + # SYSTEM_CLEANUP is what marks the close as approximated on both surfaces. + row = [r for r in session.added if isinstance(r, ExecutorPerformanceSnapshot)][0] + assert row.is_terminal is True + assert row.net_pnl_quote == Decimal("0") + + @pytest.mark.asyncio + async def test_nothing_orphaned_writes_nothing(self): + from database.repositories.executor_repository import ExecutorRepository + + session = _RecordingSession(results=[_scalars([])]) + + assert await ExecutorRepository(session).cleanup_orphaned_executors(active_executor_ids=[]) == 0 + assert session.flushed == 0 + assert session.added == [] + + +# -------------------------------------------------------------------------------------- +# The normalized row: both subjects, one shape +# -------------------------------------------------------------------------------------- + +class TestTheNormalizedRow: + def test_a_controller_row_keeps_its_whole_report(self): + """The normalization is additive, never lossy: everything with no executor + counterpart stays reachable in the passthrough.""" + row = controller_row_to_performance_row({ + "timestamp": NOW.isoformat(), + "bot_name": "bot-a", + "controller_id": "ctrl-1", + "status": "running", + "performance": { + "realized_pnl_quote": 3.0, + "unrealized_pnl_quote": 1.5, + "global_pnl_quote": 4.5, + "global_pnl_pct": 0.02, + "volume_traded": 1000.0, + "inventory_imbalance": -0.3, + "close_type_counts": {"TAKE_PROFIT": 2}, + }, + "custom_info": {"levels": 4}, + }) + + assert row.subject == "controller" + assert row.scope_id == "ctrl-1" + assert row.bot_name == "bot-a" + assert (row.realized_pnl_quote, row.unrealized_pnl_quote) == (3.0, 1.5) + assert row.volume_quote == 1000.0 + assert row.performance["inventory_imbalance"] == -0.3 + assert row.custom_info == {"levels": 4} + + def test_a_controller_reports_unknown_fees_not_zero_fees(self): + """PerformanceReport genuinely has no fees field. Zero and unknown are different + -- a consumer charting fees has to be able to tell.""" + row = controller_row_to_performance_row({ + "timestamp": NOW.isoformat(), "controller_id": "c", "status": "running", + "performance": {}, "custom_info": {}, + }) + + assert row.cum_fees_quote is None + + def test_a_live_executors_pnl_is_unrealized(self): + row = executor_row_to_performance_row( + ExecutorPerformanceRepository._to_dict(_snapshot(net_pnl="12.5", filled="4200", fees="1.8")) + ) + + assert row.subject == "executor" + assert row.scope_id == "e-1" + assert row.unrealized_pnl_quote == 12.5 + assert row.realized_pnl_quote == 0.0 + assert row.global_pnl_quote == 12.5 + assert row.volume_quote == 4200.0 + assert row.cum_fees_quote == 1.8 + assert row.is_terminal is False + + def test_a_settled_executors_pnl_is_realized(self): + row = executor_row_to_performance_row( + ExecutorPerformanceRepository._to_dict( + _snapshot(is_terminal=True, close_type="TAKE_PROFIT", net_pnl="12.5", + status="TERMINATED") + ) + ) + + assert row.realized_pnl_quote == 12.5 + assert row.unrealized_pnl_quote == 0.0 + assert row.is_terminal is True + assert row.close_type == "TAKE_PROFIT" + + def test_a_held_position_is_not_counted_as_realized(self): + """POSITION_HOLD hands the position on to position_holds; counting it realized + here would double-count it -- the same exclusion get_performance_report applies.""" + row = executor_row_to_performance_row( + ExecutorPerformanceRepository._to_dict( + _snapshot(is_terminal=True, close_type="POSITION_HOLD", net_pnl="12.5") + ) + ) + + assert row.realized_pnl_quote == 0.0 + assert row.unrealized_pnl_quote == 12.5 + + def test_an_executor_row_carries_no_heavy_passthrough(self): + row = executor_row_to_performance_row( + ExecutorPerformanceRepository._to_dict(_snapshot()) + ) + + assert row.performance == {} + assert row.custom_info == {} + + def test_both_subjects_produce_the_same_field_set(self): + """The whole point: a client writes seriesFor(scope) once.""" + controller = controller_row_to_performance_row({ + "timestamp": NOW.isoformat(), "controller_id": "c", "status": "running", + "performance": {}, "custom_info": {}, + }) + executor = executor_row_to_performance_row( + ExecutorPerformanceRepository._to_dict(_snapshot()) + ) + + assert set(controller.model_dump()) == set(executor.model_dump()) + + +# -------------------------------------------------------------------------------------- +# The route: one URL, one envelope, and filters that belong to their subject +# -------------------------------------------------------------------------------------- + +def _client(controller_history=None, executor_history=None, + controller_latest=None, executor_latest=None): + """The performance router alone, with both services stubbed.""" + from fastapi import FastAPI + from fastapi.testclient import TestClient + + import routers.performance as performance_router + from deps import get_bots_orchestrator, get_executor_service + + bots_manager = MagicMock() + executor_service = MagicMock() + + async def _controllers(**kwargs): + bots_manager.last_call = kwargs + return controller_history or ([], None, False) + + async def _executors(**kwargs): + executor_service.last_call = kwargs + return executor_history or ([], None, False) + + async def _controllers_latest(**kwargs): + bots_manager.last_latest_call = kwargs + return list(controller_latest or []) + + async def _executors_latest(**kwargs): + executor_service.last_latest_call = kwargs + return list(executor_latest or []) + + bots_manager.get_controller_performance_history = _controllers + executor_service.get_executor_performance_history = _executors + bots_manager.get_latest_controller_performance = _controllers_latest + executor_service.get_latest_executor_performance = _executors_latest + + app = FastAPI() + app.include_router(performance_router.router) + app.dependency_overrides[get_bots_orchestrator] = lambda: bots_manager + app.dependency_overrides[get_executor_service] = lambda: executor_service + return TestClient(app), bots_manager, executor_service + + +class TestTheRoute: + def test_it_serves_the_executor_series_in_the_normalized_shape(self): + rows = [ExecutorPerformanceRepository._to_dict(_snapshot(net_pnl="12.5", filled="4200"))] + client, _, _ = _client(executor_history=(rows, "2026-09-01T11:59:00+00:00", True)) + + response = client.get("/performance/history", params={"subject": "executor", "executor_id": "e-1"}) + + assert response.status_code == 200 + body = response.json() + assert body["status"] == "success" + assert body["data"][0]["subject"] == "executor" + assert body["data"][0]["scope_id"] == "e-1" + assert body["data"][0]["volume_quote"] == 4200.0 + assert body["pagination"] == { + "next_cursor": "2026-09-01T11:59:00+00:00", + "has_more": True, + "limit": 100, + "interval": "5m", + } + + def test_it_serves_the_controller_series_through_the_existing_query_path(self): + rows = [{ + "timestamp": NOW.isoformat(), "bot_name": "bot-a", "controller_id": "c-1", + "status": "running", + "performance": {"global_pnl_quote": 4.5, "volume_traded": 1000.0}, + "custom_info": {"levels": 4}, + }] + client, bots_manager, _ = _client(controller_history=(rows, None, False)) + + response = client.get("/performance/history", params={"subject": "controller", "bot_name": "bot-a"}) + + assert response.status_code == 200 + row = response.json()["data"][0] + assert row["subject"] == "controller" + assert row["global_pnl_quote"] == 4.5 + assert row["cum_fees_quote"] is None + # Nothing dropped: the raw payloads are still there. + assert row["performance"]["volume_traded"] == 1000.0 + assert row["custom_info"] == {"levels": 4} + assert bots_manager.last_call["bot_name"] == "bot-a" + + def test_an_executor_filter_on_the_controller_subject_is_a_400(self): + """FastAPI cannot express this, so it is explicit: a filter aimed at the wrong + population would otherwise be accepted and silently ignored, which reads as an + empty result rather than a mistake.""" + client, _, _ = _client() + + response = client.get("/performance/history", params={"subject": "controller", "executor_id": "e-1"}) + + assert response.status_code == 400 + assert "executor_id" in response.json()["detail"] + + def test_a_bot_name_on_the_executor_subject_is_a_400(self): + """In-process executors have no bot; it would always match nothing.""" + client, _, _ = _client() + + response = client.get("/performance/history", params={"subject": "executor", "bot_name": "bot-a"}) + + assert response.status_code == 400 + assert "bot_name" in response.json()["detail"] + + def test_controller_id_is_legal_on_both_subjects(self): + """It filters WITHIN a subject. The two namespaces are not the same thing, which + is exactly why it is never a key to join them on.""" + client, _, _ = _client() + + for subject in ("controller", "executor"): + assert client.get( + "/performance/history", params={"subject": subject, "controller_id": "main"} + ).status_code == 200 + + def test_the_subject_is_required(self): + client, _, _ = _client() + assert client.get("/performance/history").status_code == 422 + + def test_the_limit_is_clamped_at_a_thousand(self): + client, _, _ = _client() + assert client.get( + "/performance/history", params={"subject": "executor", "limit": 1001} + ).status_code == 422 + + def test_one_minute_is_only_meaningful_on_the_executor_subject_but_accepted_on_both(self): + """`interval` is a floor: asking a 5-minute controller series for 1m returns its + native grain rather than an error.""" + client, _, _ = _client() + assert client.get( + "/performance/history", params={"subject": "controller", "interval": "1m"} + ).status_code == 200 + + def test_a_malformed_timestamp_is_a_400_not_a_500(self): + client, _, _ = _client() + + response = client.get( + "/performance/history", params={"subject": "executor", "start_time": "yesterday"} + ) + + assert response.status_code == 400 + assert "Invalid datetime format" in response.json()["detail"] + + +# -------------------------------------------------------------------------------------- +# The existing controller routes are untouched +# -------------------------------------------------------------------------------------- + +class TestTheControllerPathIsUntouched: + def test_the_controller_repository_still_assumes_its_own_five_minute_grain(self): + """It is on a wire-compatible path. The new repository re-implements the sampler + with a grain parameter rather than generalizing this one.""" + import inspect + + from database.repositories.controller_performance_repository import ControllerPerformanceRepository + + source = inspect.getsource(ControllerPerformanceRepository) + assert "interval_minutes <= 5" in source + assert "interval_minutes // 5" in source + assert "grain_minutes" not in source + + def test_the_unified_route_reuses_the_controller_query_path(self): + """Sharing one call is what stops the two routes from ever drifting.""" + import inspect + + import routers.performance as performance_router + + source = inspect.getsource(performance_router) + assert "get_controller_performance_history" in source + + def test_the_existing_controller_routes_still_answer_the_old_shape(self): + import inspect + + import routers.bot_orchestration as bot_orchestration + + source = inspect.getsource(bot_orchestration) + assert '@router.get("/controller-performance-latest")' in source + assert '@router.get("/controller-performance-history")' in source + # Nothing from the new normalization leaked onto the old path. + assert "PerformanceRow" not in source + + +# -------------------------------------------------------------------------------------- +# Helpers used above +# -------------------------------------------------------------------------------------- + +def _metadata(_executor_id): + return { + "executor_type": "position_executor", + "account_name": "master_account", + "connector_name": "binance_perpetual", + "trading_pair": "BTC-USDT", + "controller_id": "main", + "created_at": NOW, + } + + +def _with_session(service, session): + """Point the service's db_manager at a fixed fake session.""" + from contextlib import asynccontextmanager + + @asynccontextmanager + async def _context(): + yield session + + service.db_manager = SimpleNamespace(get_session_context=_context) + + +# -------------------------------------------------------------------------------------- +# /performance/latest: the current value of every scope, so a consumer can drop +# /bot-orchestration/controller-performance-latest as well as -history +# -------------------------------------------------------------------------------------- + +def _controller_latest_row(controller_id="c-1", bot_name="bot-a", timestamp=NOW, pnl=4.5): + return { + "timestamp": timestamp.isoformat(), + "bot_name": bot_name, + "controller_id": controller_id, + "status": "running", + "performance": {"global_pnl_quote": pnl, "volume_traded": 1000.0}, + "custom_info": {"levels": 4}, + } + + +class TestTheLatestQuery: + @pytest.mark.asyncio + async def test_it_returns_the_last_row_of_each_executor(self): + session = _RecordingSession([_scalars([ + _snapshot(executor_id="e-2", net_pnl="7"), + _snapshot(executor_id="e-1", net_pnl="3"), + ])]) + + rows = await ExecutorPerformanceRepository(session).get_latest() + + assert [r["executor_id"] for r in rows] == ["e-2", "e-1"] + assert rows[0]["net_pnl_quote"] == 7.0 + + @pytest.mark.asyncio + async def test_it_orders_newest_first_and_takes_a_limit(self): + """Every executor that ever ran leaves a terminal row, so the unfiltered result + would grow without bound. Newest-first plus a limit puts the executors that are + still being snapshotted at the top.""" + session = _RecordingSession([_scalars([])]) + + await ExecutorPerformanceRepository(session).get_latest(limit=25) + + sql = str(session.executed[0].compile(compile_kwargs={"literal_binds": True})) + assert "ORDER BY" in sql and "DESC" in sql + assert "LIMIT 25" in sql + + @pytest.mark.asyncio + async def test_no_limit_means_no_limit_clause(self): + session = _RecordingSession([_scalars([])]) + + await ExecutorPerformanceRepository(session).get_latest() + + assert "LIMIT" not in str(session.executed[0].compile(compile_kwargs={"literal_binds": True})) + + @pytest.mark.asyncio + async def test_a_filter_narrows_the_grouped_subquery(self): + """The filter decides which executors are aggregated at all, rather than + aggregating the whole table and discarding most of it afterwards.""" + session = _RecordingSession([_scalars([])]) + + await ExecutorPerformanceRepository(session).get_latest(account_name="master_account") + + sql = str(session.executed[0].compile(compile_kwargs={"literal_binds": True})) + grouped = sql.split("GROUP BY")[0] + assert "master_account" in grouped, "the filter landed outside the grouped subquery" + + @pytest.mark.asyncio + async def test_a_closed_executors_last_row_is_its_terminal_row(self): + """So "the final value" needs no second call and no join to `executors`.""" + session = _RecordingSession([_scalars([ + _snapshot(is_terminal=True, close_type="TAKE_PROFIT", net_pnl="9", status="TERMINATED"), + ])]) + + rows = await ExecutorPerformanceRepository(session).get_latest() + + assert rows[0]["is_terminal"] is True + assert rows[0]["close_type"] == "TAKE_PROFIT" + + +class TestTheLatestRoute: + def test_it_serves_the_executor_scopes_in_the_normalized_shape(self): + rows = [ExecutorPerformanceRepository._to_dict(_snapshot(net_pnl="12.5", filled="4200"))] + client, _, _ = _client(executor_latest=rows) + + response = client.get("/performance/latest", params={"subject": "executor"}) + + assert response.status_code == 200 + body = response.json() + assert body["status"] == "success" + assert body["data"][0]["subject"] == "executor" + assert body["data"][0]["scope_id"] == "e-1" + assert body["data"][0]["volume_quote"] == 4200.0 + # One row per scope is not a series: there is nothing to cursor through. + assert "pagination" not in body + + def test_it_serves_the_controller_scopes_through_the_existing_query_path(self): + client, bots_manager, _ = _client(controller_latest=[_controller_latest_row()]) + + response = client.get("/performance/latest", params={"subject": "controller", "bot_name": "bot-a"}) + + assert response.status_code == 200 + row = response.json()["data"][0] + assert row["subject"] == "controller" + assert row["scope_id"] == "c-1" + assert row["global_pnl_quote"] == 4.5 + assert row["cum_fees_quote"] is None + assert row["performance"]["volume_traded"] == 1000.0 + assert bots_manager.last_latest_call["bot_name"] == "bot-a" + + def test_the_controller_scopes_come_back_newest_first(self): + """get_latest_controller_performance returns join order, so the route sorts -- + otherwise `limit` would truncate an arbitrary set of scopes.""" + older = _controller_latest_row(controller_id="old", timestamp=NOW - timedelta(hours=2)) + newer = _controller_latest_row(controller_id="new", timestamp=NOW) + client, _, _ = _client(controller_latest=[older, newer]) + + response = client.get("/performance/latest", params={"subject": "controller"}) + + assert [r["scope_id"] for r in response.json()["data"]] == ["new", "old"] + + def test_a_malformed_controller_timestamp_sorts_last_instead_of_500ing(self): + """One bad row must not take down a dashboard's whole tile set.""" + bad = _controller_latest_row(controller_id="bad") + bad["timestamp"] = "not-a-timestamp" + client, _, _ = _client(controller_latest=[bad, _controller_latest_row(controller_id="good")]) + + response = client.get("/performance/latest", params={"subject": "controller"}) + + assert response.status_code == 200 + assert [r["scope_id"] for r in response.json()["data"]] == ["good", "bad"] + + def test_controller_id_narrows_the_controller_scopes(self): + """get_latest_controller_performance only takes bot_name -- it is the method the + wire-compatible route calls and is deliberately unchanged -- so the route filters.""" + client, _, _ = _client(controller_latest=[ + _controller_latest_row(controller_id="c-1"), + _controller_latest_row(controller_id="c-2"), + ]) + + response = client.get("/performance/latest", + params={"subject": "controller", "controller_id": "c-2"}) + + assert [r["scope_id"] for r in response.json()["data"]] == ["c-2"] + + def test_the_limit_caps_the_controller_scopes(self): + client, _, _ = _client(controller_latest=[ + _controller_latest_row(controller_id=f"c-{i}", timestamp=NOW - timedelta(minutes=i)) + for i in range(5) + ]) + + response = client.get("/performance/latest", params={"subject": "controller", "limit": 2}) + + assert [r["scope_id"] for r in response.json()["data"]] == ["c-0", "c-1"] + + def test_the_limit_reaches_the_executor_query(self): + client, _, executor_service = _client() + + client.get("/performance/latest", params={"subject": "executor", "limit": 7}) + + assert executor_service.last_latest_call["limit"] == 7 + + def test_an_executor_filter_on_the_controller_subject_is_a_400(self): + client, _, _ = _client() + + response = client.get("/performance/latest", + params={"subject": "controller", "executor_id": "e-1"}) + + assert response.status_code == 400 + assert "executor_id" in response.json()["detail"] + + def test_a_bot_name_on_the_executor_subject_is_a_400(self): + client, _, _ = _client() + + response = client.get("/performance/latest", + params={"subject": "executor", "bot_name": "bot-a"}) + + assert response.status_code == 400 + assert "bot_name" in response.json()["detail"] + + def test_both_routes_enforce_the_same_filter_rule(self): + """One helper, so /history and /latest cannot drift into disagreeing about which + filter belongs to which population.""" + client, _, _ = _client() + + for path in ("/performance/history", "/performance/latest"): + assert client.get(path, params={"subject": "controller", "trading_pair": "BTC-USDT"}).status_code == 400 + assert client.get(path, params={"subject": "executor", "bot_name": "b"}).status_code == 400 + + def test_the_subject_is_required(self): + client, _, _ = _client() + + assert client.get("/performance/latest").status_code == 422 + + def test_the_limit_is_clamped_at_a_thousand(self): + client, _, _ = _client() + + assert client.get("/performance/latest", + params={"subject": "executor", "limit": 1001}).status_code == 422 + + def test_both_subjects_produce_the_same_field_set(self): + """The whole point: latestFor(scope) is written once, not twice.""" + executor_client, _, _ = _client( + executor_latest=[ExecutorPerformanceRepository._to_dict(_snapshot())]) + controller_client, _, _ = _client(controller_latest=[_controller_latest_row()]) + + executor_row = executor_client.get( + "/performance/latest", params={"subject": "executor"}).json()["data"][0] + controller_row = controller_client.get( + "/performance/latest", params={"subject": "controller"}).json()["data"][0] + + assert executor_row.keys() == controller_row.keys() + + def test_it_shares_the_row_shape_with_the_history_route(self): + """A dashboard's live tiles and its charts read the same fields off one client.""" + row = ExecutorPerformanceRepository._to_dict(_snapshot()) + client, _, _ = _client(executor_latest=[row], executor_history=([row], None, False)) + + latest = client.get("/performance/latest", params={"subject": "executor"}).json()["data"][0] + history = client.get("/performance/history", params={"subject": "executor"}).json()["data"][0] + + assert latest == history + + def test_the_old_latest_route_still_exists(self): + """Wire compatibility is absolute: this is new surface, and a consumer migrates + when it chooses to.""" + import inspect + + import routers.bot_orchestration as bot_orchestration + + source = inspect.getsource(bot_orchestration) + assert '@router.get("/controller-performance-latest")' in source diff --git a/test/test_executor_stats_is_two_queries.py b/test/test_executor_stats_is_two_queries.py new file mode 100644 index 00000000..5ce4150e --- /dev/null +++ b/test/test_executor_stats_is_two_queries.py @@ -0,0 +1,226 @@ +"""`get_executor_stats` asks the database twice, not seven times. + +The method used to issue four unfiltered scalar queries -- two COUNTs and two SUMs, none +of them narrowing anything -- followed by three separate GROUP BY queries over the same +table. Seven round-trips for a payload that is entirely aggregates: every one of them a +full pass over the executors table, and the four scalars share a single filter (none). + +It is now one aggregate row plus one grouped statement over +(executor_type, status, connector_name), pivoted back into the three breakdowns in Python +-- the same shape `get_performance_report` was reduced to. What is pinned here is the +round-trip count, and that the dict handed back is the one the seven queries produced. +""" + +from contextlib import asynccontextmanager +from decimal import Decimal +from types import SimpleNamespace + +import pytest +from sqlalchemy import create_engine, func, select +from sqlalchemy.orm import Session +from sqlalchemy.pool import StaticPool + +from database.models import ExecutorRecord +from database.repositories.executor_repository import ExecutorRepository + + +class _AsyncSessionAdapter: + """The async surface the repository uses, over a real synchronous Session. + + aiosqlite is not installed here and the method only ever awaits `execute`, so this + runs the real SQL against the real schema, and counts the statements it is asked for. + """ + + def __init__(self, session: Session): + self._session = session + self.statements = [] + + async def execute(self, statement): + self.statements.append(statement) + return self._session.execute(statement) + + async def commit(self): + self._session.commit() + + async def rollback(self): + self._session.rollback() + + async def close(self): + self._session.close() + + +@pytest.fixture +def db(): + engine = create_engine("sqlite://", connect_args={"check_same_thread": False}, + poolclass=StaticPool) + ExecutorRecord.__table__.create(engine) + + @asynccontextmanager + async def session_context(): + session = Session(engine) + adapter = _AsyncSessionAdapter(session) + try: + yield adapter + await adapter.commit() + finally: + await adapter.close() + + def insert(rows): + """rows: dicts overriding the defaults below, one executor each.""" + with Session(engine) as session: + for i, row in enumerate(rows): + fields = { + "executor_type": "position_executor", + "connector_name": "binance_perpetual", + "status": "TERMINATED", + "net_pnl_quote": Decimal("1"), + "filled_amount_quote": Decimal("100"), + **row, + } + session.add(ExecutorRecord( + executor_id=f"e-{i}", + account_name="master", + trading_pair="BTC-USDT", + controller_id="main", + close_type="TAKE_PROFIT", + net_pnl_pct=Decimal("0"), + cum_fees_quote=Decimal("0"), + **fields, + )) + session.commit() + + try: + yield SimpleNamespace(session_context=session_context, insert=insert, engine=engine) + finally: + engine.dispose() + + +def _stats_the_old_way(engine): + """The seven queries the method used to run, verbatim, as the expected result.""" + with Session(engine) as session: + total = session.execute(select(func.count(ExecutorRecord.id))).scalar() or 0 + active = session.execute( + select(func.count(ExecutorRecord.id)).where(ExecutorRecord.status == "RUNNING") + ).scalar() or 0 + pnl = session.execute(select(func.sum(ExecutorRecord.net_pnl_quote))).scalar() or Decimal("0") + volume = session.execute( + select(func.sum(ExecutorRecord.filled_amount_quote)) + ).scalar() or Decimal("0") + + def grouped(column): + rows = session.execute( + select(column, func.count(ExecutorRecord.id).label("count")).group_by(column) + ) + return {row[0]: row.count for row in rows} + + return { + "total_executors": total, + "active_executors": active, + "total_pnl_quote": float(pnl), + "total_volume_quote": float(volume), + "type_counts": grouped(ExecutorRecord.executor_type), + "status_counts": grouped(ExecutorRecord.status), + "connector_counts": grouped(ExecutorRecord.connector_name), + } + + +async def _stats(db): + """Runs the method and reports both its answer and how many statements it cost.""" + async with db.session_context() as session: + stats = await ExecutorRepository(session).get_executor_stats() + return stats, len(session.statements) + + +_A_MIXED_TABLE = [ + {"status": "RUNNING", "executor_type": "position_executor", "connector_name": "binance_perpetual"}, + {"status": "RUNNING", "executor_type": "dca_executor", "connector_name": "kucoin"}, + {"status": "TERMINATED", "executor_type": "position_executor", "connector_name": "kucoin", + "net_pnl_quote": Decimal("-3.5"), "filled_amount_quote": Decimal("250.25")}, + {"status": "FAILED", "executor_type": "arbitrage_executor", "connector_name": "binance", + "net_pnl_quote": None}, + {"status": "TERMINATED", "executor_type": "dca_executor", "connector_name": "binance", + "net_pnl_quote": Decimal("12.75"), "filled_amount_quote": Decimal("0")}, +] + + +# -------------------------------------------------------------------------------------- +# The round-trip count +# -------------------------------------------------------------------------------------- + +@pytest.mark.asyncio +async def test_the_stats_cost_two_queries(db): + db.insert(_A_MIXED_TABLE) + + _, queries = await _stats(db) + + assert queries == 2, ( + f"get_executor_stats issued {queries} statements; it is one aggregate row plus " + "one grouped statement, and every extra one is another pass over the table" + ) + + +@pytest.mark.asyncio +async def test_an_empty_table_costs_the_same_two_queries(db): + _, queries = await _stats(db) + + assert queries == 2 + + +@pytest.mark.asyncio +async def test_the_query_count_does_not_follow_the_number_of_groups(db): + """Ten distinct connectors, still two statements -- the breakdowns are pivoted here.""" + db.insert([{"connector_name": f"venue_{i}", "executor_type": f"type_{i % 3}"} + for i in range(10)]) + + stats, queries = await _stats(db) + + assert queries == 2 + assert len(stats["connector_counts"]) == 10 + + +# -------------------------------------------------------------------------------------- +# The dict itself did not change +# -------------------------------------------------------------------------------------- + +@pytest.mark.asyncio +@pytest.mark.parametrize("rows", [ + pytest.param([], id="empty_table"), + pytest.param(_A_MIXED_TABLE, id="mixed_types_statuses_and_connectors"), + pytest.param([{"status": "RUNNING"}] * 3, id="all_running"), + pytest.param([{"net_pnl_quote": None, "filled_amount_quote": None}], id="null_amounts"), + pytest.param([{"net_pnl_quote": Decimal("-1.5")}, {"net_pnl_quote": Decimal("1.5")}], + id="pnl_cancels_to_zero"), +]) +async def test_the_result_matches_the_seven_query_version(db, rows): + db.insert(rows) + + stats, _ = await _stats(db) + + assert stats == _stats_the_old_way(db.engine) + + +@pytest.mark.asyncio +async def test_the_keys_are_still_the_ones_the_endpoint_returns(db): + db.insert(_A_MIXED_TABLE) + + stats, _ = await _stats(db) + + assert set(stats) == { + "total_executors", "active_executors", "total_pnl_quote", "total_volume_quote", + "type_counts", "status_counts", "connector_counts", + } + assert stats["total_executors"] == 5 + assert stats["active_executors"] == 2 + assert isinstance(stats["total_pnl_quote"], float) + assert isinstance(stats["total_volume_quote"], float) + + +@pytest.mark.asyncio +async def test_a_status_with_no_executors_is_absent_rather_than_zero(db): + """The grouped pivot must not invent keys the GROUP BY never produced.""" + db.insert([{"status": "RUNNING"}, {"status": "RUNNING"}]) + + stats, _ = await _stats(db) + + assert stats["status_counts"] == {"RUNNING": 2} + assert stats["active_executors"] == 2 diff --git a/test/test_executor_volume_is_the_filled_amount.py b/test/test_executor_volume_is_the_filled_amount.py index 6daa41d4..dbf18d0b 100644 --- a/test/test_executor_volume_is_the_filled_amount.py +++ b/test/test_executor_volume_is_the_filled_amount.py @@ -21,7 +21,7 @@ import re from unittest.mock import MagicMock -from database.models import ExecutorRecord +from database.models import ExecutorPerformanceSnapshot, ExecutorRecord from database.repositories.executor_repository import ExecutorRepository from services.executor_service import ExecutorService @@ -32,6 +32,31 @@ def test_the_record_has_no_separate_volume_column(): ) +def test_the_performance_snapshot_has_no_separate_volume_column(): + """The same rule on the table FEAT-001 added, which mirrors these four metrics.""" + assert not hasattr(ExecutorPerformanceSnapshot, "volume_traded_quote"), ( + "the snapshot table grew a second volume column; filled_amount_quote is it" + ) + + +def test_the_normalized_performance_row_reads_the_filled_amount_as_volume(): + """The unified /performance/history route maps volume_quote from the one field. + + A row that read a second column here would put the LP deposit back into volume by the + back door, on the newest reader rather than the oldest. + """ + from models.performance import executor_row_to_performance_row + + row = executor_row_to_performance_row({ + "timestamp": "2026-09-01T12:00:00+00:00", + "executor_id": "e-1", + "status": "RUNNING", + "filled_amount_quote": 2500.0, + }) + + assert row.volume_quote == 2500.0 + + def test_every_aggregate_sums_the_filled_amount(): """Three aggregates summed the old column; a fourth added later would too.""" source = inspect.getsource(ExecutorRepository.get_performance_report) diff --git a/test/test_executor_ws_push_loops.py b/test/test_executor_ws_push_loops.py new file mode 100644 index 00000000..92f11694 --- /dev/null +++ b/test/test_executor_ws_push_loops.py @@ -0,0 +1,566 @@ +""" +Tests for the shared /ws/executors push loop (ARCH-055). + +Seven near-identical `_*_push_loop` coroutines were collapsed into one generic +`_push_loop` driven by a `sub_type -> PushSpec(fetch, msg_type, extra)` table, +with `_logs_push_loop` kept separate because it keys on `last_log_count`. +These tests pin the behaviour that must not have changed: the frame shape of +every channel, per-subscription intervals, change detection, per-channel error +handling, and cancellation/cleanup. + +Run with: pytest test/test_executor_ws_push_loops.py -v --asyncio-mode=auto +""" +import asyncio +from datetime import datetime +from unittest.mock import AsyncMock, MagicMock + +import pytest +from fastapi.websockets import WebSocketDisconnect + +from services.executor_ws_manager import SUBSCRIPTION_TYPES, ExecutorSubscription, ExecutorWebSocketManager + + +class RecordingWebSocket: + """Captures pushed frames and lets a test await the Nth one.""" + + def __init__(self, raise_on_send=None): + self.sent = [] + self.raise_on_send = raise_on_send + self._target = None + self._reached = asyncio.Event() + + async def send_json(self, message): + if self.raise_on_send is not None: + raise self.raise_on_send + self.sent.append(message) + if self._target is not None and len(self.sent) >= self._target: + self._reached.set() + + async def wait_for_frames(self, count, timeout=2.0): + self._target = count + if len(self.sent) >= count: + return + self._reached.clear() + await asyncio.wait_for(self._reached.wait(), timeout) + + def data_frames(self): + return [f for f in self.sent if f.get("type") not in ("subscribed", "unsubscribed", "error")] + + +class FakePosition: + """Minimal stand-in for a PositionHeld row, with real numbers to cast.""" + + def __init__(self, trading_pair="BTC-USDT"): + self.trading_pair = trading_pair + self.connector_name = "binance_perpetual" + self.account_name = "master" + self.controller_id = "ctrl-1" + self.buy_amount_base = 1.0 + self.buy_amount_quote = 100.0 + self.sell_amount_base = 0.5 + self.sell_amount_quote = 60.0 + self.net_amount_base = 0.5 + self.buy_breakeven_price = 100.0 + self.sell_breakeven_price = 120.0 + self.matched_amount_base = 0.5 + self.unmatched_amount_base = 0.5 + self.position_side = "LONG" + self.realized_pnl_quote = 10.0 + self.executor_ids = ["e-1", "e-2"] + self.last_updated = datetime(2026, 1, 1, 12, 0, 0) + + def get_unrealized_pnl(self, rate): + return float(rate) * self.net_amount_base + + +def build_manager(**overrides): + """A manager whose every backing service call returns canned data.""" + executor_service = MagicMock() + executor_service.get_executors = AsyncMock(return_value=[{"id": "e-1"}]) + executor_service.get_executor = AsyncMock(return_value={"id": "e-1"}) + executor_service.get_summary = MagicMock(return_value={"total": 1}) + executor_service.get_performance_report = AsyncMock(return_value={"pnl": 1.0}) + executor_service.get_positions_held = MagicMock(return_value=[FakePosition()]) + executor_service.get_executor_logs = MagicMock(return_value=[{"msg": "hello"}]) + + market_data_service = MagicMock() + market_data_service.get_rate = MagicMock(return_value=200.0) + + orchestrator = MagicMock() + orchestrator.get_bot_status = MagicMock( + return_value={ + "status": "running", + "performance": {"pnl": 1.0}, + "recently_active": True, + "logs": ["should be stripped"], + } + ) + orchestrator.get_all_bots_status = MagicMock( + return_value={ + "bot-1": { + "status": "running", + "source": "broker", + "performance": {"pnl": 1.0}, + "recently_active": True, + "logs": ["should be stripped"], + } + } + ) + + for name, value in overrides.items(): + setattr(executor_service, name, value) + + return ExecutorWebSocketManager( + executor_service=executor_service, + market_data_service=market_data_service, + bots_orchestrator=orchestrator, + ), executor_service, orchestrator + + +SUBSCRIBE_MESSAGES = { + "executors": {"type": "executors"}, + "executor_detail": {"type": "executor_detail", "executor_id": "e-1"}, + "executor_summary": {"type": "executor_summary"}, + "performance": {"type": "performance", "controller_id": "ctrl-1"}, + "positions": {"type": "positions"}, + "executor_logs": {"type": "executor_logs", "executor_id": "e-1"}, + "bot_status": {"type": "bot_status", "bot_name": "bot-1"}, + "all_bots_status": {"type": "all_bots_status"}, +} + +EXPECTED_FRAME_KEYS = { + "executors": ["type", "subscription_id", "data", "total_count", "timestamp"], + "executor_detail": ["type", "subscription_id", "data", "timestamp"], + "executor_summary": ["type", "subscription_id", "data", "timestamp"], + "performance": ["type", "subscription_id", "data", "timestamp"], + "positions": ["type", "subscription_id", "data", "timestamp"], + "executor_logs": ["type", "subscription_id", "data", "total_count", "timestamp"], + "bot_status": ["type", "subscription_id", "data", "timestamp"], + "all_bots_status": ["type", "subscription_id", "data", "bot_count", "timestamp"], +} + + +async def subscribe_and_collect(manager, websocket, sub_type, frames=2): + """Subscribe, wait for the ack plus the first data frame, then clean up.""" + await manager.handle_subscribe("conn-1", websocket, SUBSCRIBE_MESSAGES[sub_type]) + try: + await websocket.wait_for_frames(frames) + finally: + manager.remove_connection("conn-1") + await asyncio.sleep(0) + + +class TestDispatchTable: + """_get_push_fn must still resolve every declared subscription type.""" + + def test_every_subscription_type_resolves_to_a_loop(self): + manager, _, _ = build_manager() + for sub_type in SUBSCRIPTION_TYPES: + assert manager._get_push_fn(sub_type) is not None + + def test_the_spec_table_covers_every_type_except_logs(self): + manager, _, _ = build_manager() + assert set(manager._push_specs()) == SUBSCRIPTION_TYPES - {"executor_logs"} + + def test_each_spec_message_type_matches_its_subscription_type(self): + manager, _, _ = build_manager() + for sub_type, spec in manager._push_specs().items(): + assert spec.msg_type == sub_type + + def test_only_two_polling_loop_bodies_remain(self): + """The seven hash-and-send copies are gone (ARCH-055 acceptance).""" + from pathlib import Path + + import services.executor_ws_manager as module + + source = Path(module.__file__).read_text() + assert source.count("await asyncio.sleep(sub.update_interval)") == 2 + assert source.count("while True:") == 2 + + +class TestFrameShapePerChannel: + """Each of the eight channels still emits its original frame shape.""" + + @pytest.mark.parametrize("sub_type", sorted(SUBSCRIPTION_TYPES)) + async def test_ack_then_a_frame_of_the_subscription_type(self, sub_type): + manager, _, _ = build_manager() + websocket = RecordingWebSocket() + + await subscribe_and_collect(manager, websocket, sub_type) + + assert websocket.sent[0]["type"] == "subscribed" + assert websocket.sent[0]["subscription_type"] == sub_type + data_frames = websocket.data_frames() + assert data_frames, f"{sub_type} pushed no data frame" + assert data_frames[0]["type"] == sub_type + + @pytest.mark.parametrize("sub_type", sorted(SUBSCRIPTION_TYPES)) + async def test_frame_keys_and_their_order_are_unchanged(self, sub_type): + manager, _, _ = build_manager() + websocket = RecordingWebSocket() + + await subscribe_and_collect(manager, websocket, sub_type) + + frame = websocket.data_frames()[0] + assert list(frame) == EXPECTED_FRAME_KEYS[sub_type] + assert frame["subscription_id"] == websocket.sent[0]["subscription_id"] + assert isinstance(frame["timestamp"], float) + + async def test_executors_frame_carries_the_total_count(self): + manager, _, _ = build_manager() + websocket = RecordingWebSocket() + + await subscribe_and_collect(manager, websocket, "executors") + + frame = websocket.data_frames()[0] + assert frame["data"] == [{"id": "e-1"}] + assert frame["total_count"] == 1 + + async def test_all_bots_status_frame_carries_the_bot_count_and_strips_logs(self): + manager, _, _ = build_manager() + websocket = RecordingWebSocket() + + await subscribe_and_collect(manager, websocket, "all_bots_status") + + frame = websocket.data_frames()[0] + assert frame["bot_count"] == 1 + assert frame["data"] == { + "bot-1": { + "status": "running", + "source": "broker", + "performance": {"pnl": 1.0}, + "recently_active": True, + } + } + + async def test_bot_status_frame_strips_logs(self): + manager, _, _ = build_manager() + websocket = RecordingWebSocket() + + await subscribe_and_collect(manager, websocket, "bot_status") + + assert websocket.data_frames()[0]["data"] == { + "bot_name": "bot-1", + "status": "running", + "performance": {"pnl": 1.0}, + "recently_active": True, + } + + async def test_positions_payload_shaping_is_preserved(self): + manager, _, _ = build_manager() + websocket = RecordingWebSocket() + + await subscribe_and_collect(manager, websocket, "positions") + + payload = websocket.data_frames()[0]["data"] + assert payload["total_positions"] == 1 + assert payload["total_realized_pnl"] == 10.0 + assert payload["total_unrealized_pnl"] == 100.0 + position = payload["positions"][0] + assert list(position) == [ + "trading_pair", "connector_name", "account_name", "controller_id", + "buy_amount_base", "buy_amount_quote", "sell_amount_base", + "sell_amount_quote", "net_amount_base", "buy_breakeven_price", + "sell_breakeven_price", "matched_amount_base", "unmatched_amount_base", + "position_side", "realized_pnl_quote", "unrealized_pnl_quote", + "executor_count", "executor_ids", "last_updated", + ] + assert position["unrealized_pnl_quote"] == 100.0 + assert position["executor_count"] == 2 + assert position["last_updated"] == "2026-01-01T12:00:00" + + async def test_positions_without_a_rate_report_no_unrealized_pnl(self): + manager, _, _ = build_manager() + manager._market_data_service.get_rate = MagicMock(return_value=None) + websocket = RecordingWebSocket() + + await subscribe_and_collect(manager, websocket, "positions") + + payload = websocket.data_frames()[0]["data"] + assert payload["total_unrealized_pnl"] is None + assert payload["positions"][0]["unrealized_pnl_quote"] is None + + async def test_logs_frame_sends_only_the_new_entries(self): + manager, executor_service, _ = build_manager() + executor_service.get_executor_logs = MagicMock( + side_effect=[[{"n": 1}], [{"n": 1}, {"n": 2}]] + ) + websocket = RecordingWebSocket() + sub = ExecutorSubscription( + sub_id="executor_logs_e-1", sub_type="executor_logs", + update_interval=0, executor_id="e-1", + ) + + await run_loop(manager, "executor_logs", sub, websocket, frames=2) + + assert [f["data"] for f in websocket.sent] == [[{"n": 1}], [{"n": 2}]] + assert [f["total_count"] for f in websocket.sent] == [1, 2] + + +def make_sub(sub_type, interval=0, **kwargs): + return ExecutorSubscription( + sub_id=f"{sub_type}_test", sub_type=sub_type, update_interval=interval, **kwargs + ) + + +async def run_loop(manager, sub_type, sub, websocket, frames=1, timeout=2.0): + """Drive one push loop directly until N frames land, then cancel it.""" + push_fn = manager._get_push_fn(sub_type) + task = asyncio.create_task(push_fn("conn-1", websocket, sub)) + try: + await websocket.wait_for_frames(frames, timeout) + finally: + task.cancel() + await asyncio.gather(task, return_exceptions=True) + return task + + +class TestChangeDetection: + """A repeated payload must not be re-sent; a changed one must be sent once.""" + + async def test_identical_data_is_pushed_only_once(self): + manager, executor_service, _ = build_manager() + executor_service.get_executors = AsyncMock(return_value=[{"id": "e-1"}]) + websocket = RecordingWebSocket() + sub = make_sub("executors") + + task = asyncio.create_task( + manager._get_push_fn("executors")("conn-1", websocket, sub) + ) + await websocket.wait_for_frames(1) + for _ in range(20): + await asyncio.sleep(0) + task.cancel() + await asyncio.gather(task, return_exceptions=True) + + assert executor_service.get_executors.await_count > 1 + assert len(websocket.sent) == 1 + + async def test_changed_data_is_pushed_exactly_once_more(self): + manager, executor_service, _ = build_manager() + executor_service.get_executors = AsyncMock( + side_effect=[[{"id": "e-1"}], [{"id": "e-1"}], [{"id": "e-2"}], [{"id": "e-2"}]] + ) + websocket = RecordingWebSocket() + sub = make_sub("executors") + + await run_loop(manager, "executors", sub, websocket, frames=2) + + assert [f["data"] for f in websocket.sent] == [[{"id": "e-1"}], [{"id": "e-2"}]] + + async def test_the_hash_is_kept_on_the_subscription(self): + manager, _, _ = build_manager() + websocket = RecordingWebSocket() + sub = make_sub("executor_summary") + + await run_loop(manager, "executor_summary", sub, websocket, frames=1) + + assert sub.last_sent_hash is not None + + +class TestErrorHandling: + """A failing fetch is logged per channel and the loop keeps polling. + + Since CORR-111 the generic loop also pushes one `error` frame before it + retries, so the client is told the channel is failing; `executor_logs` keeps + its own loop and still only logs. See test_performance_report_reports_db_failure.py. + """ + + @pytest.mark.parametrize( + "sub_type,attr,is_async", + [ + ("executors", "get_executors", True), + ("executor_detail", "get_executor", True), + ("executor_summary", "get_summary", False), + ("performance", "get_performance_report", True), + ("positions", "get_positions_held", False), + ("executor_logs", "get_executor_logs", False), + ], + ) + async def test_a_fetch_failure_does_not_kill_the_loop(self, sub_type, attr, is_async, caplog): + manager, executor_service, _ = build_manager() + good = { + "executors": [{"id": "e-1"}], + "executor_logs": [{"msg": "x"}], + "positions": [], + }.get(sub_type, {"ok": True}) + maker = AsyncMock if is_async else MagicMock + setattr(executor_service, attr, maker(side_effect=[RuntimeError("boom"), good, good])) + websocket = RecordingWebSocket() + sub = make_sub(sub_type, executor_id="e-1") + + # Every generic channel emits an error frame first, then the recovery frame. + expected_frames = 1 if sub_type == "executor_logs" else 2 + + with caplog.at_level("ERROR"): + await run_loop(manager, sub_type, sub, websocket, frames=expected_frames) + + assert websocket.data_frames(), "the loop did not recover after the failing poll" + assert websocket.data_frames()[0]["type"] == sub_type + assert f"[WS-Exec] {sub_type} push error" in caplog.text + + async def test_an_orchestrator_failure_is_logged_under_its_own_channel(self, caplog): + manager, _, orchestrator = build_manager() + orchestrator.get_all_bots_status = MagicMock( + side_effect=[RuntimeError("boom"), {"bot-1": {"status": "running"}}] + ) + websocket = RecordingWebSocket() + sub = make_sub("all_bots_status") + + with caplog.at_level("ERROR"): + await run_loop(manager, "all_bots_status", sub, websocket, frames=2) + + assert "[WS-Exec] all_bots_status push error" in caplog.text + assert websocket.data_frames()[0]["type"] == "all_bots_status" + + +class TestDisconnectStopsTheLoop: + """A dropped client ends the loop instead of erroring every interval.""" + + @pytest.mark.parametrize("error", [WebSocketDisconnect(), RuntimeError("closed")]) + @pytest.mark.parametrize("sub_type", sorted(SUBSCRIPTION_TYPES)) + async def test_a_disconnect_breaks_out_of_the_loop(self, sub_type, error, caplog): + manager, _, _ = build_manager() + websocket = RecordingWebSocket(raise_on_send=error) + sub = make_sub(sub_type, executor_id="e-1", bot_name="bot-1") + + with caplog.at_level("ERROR"): + task = asyncio.create_task( + manager._get_push_fn(sub_type)("conn-1", websocket, sub) + ) + await asyncio.wait_for(task, timeout=2.0) + + assert task.done() + assert "push error" not in caplog.text + + +class TestCancellationAndCleanup: + """Cancellation is swallowed and every teardown path cancels its tasks.""" + + @pytest.mark.parametrize("sub_type", sorted(SUBSCRIPTION_TYPES)) + async def test_cancelling_a_loop_raises_nothing(self, sub_type): + manager, _, _ = build_manager() + websocket = RecordingWebSocket() + sub = make_sub(sub_type, interval=5.0, executor_id="e-1", bot_name="bot-1") + + task = asyncio.create_task(manager._get_push_fn(sub_type)("conn-1", websocket, sub)) + await websocket.wait_for_frames(1) + task.cancel() + await task + + assert task.done() + assert task.exception() is None + + async def test_unsubscribe_cancels_the_task(self): + manager, _, _ = build_manager() + websocket = RecordingWebSocket() + + await manager.handle_subscribe("conn-1", websocket, SUBSCRIBE_MESSAGES["executors"]) + sub_id = websocket.sent[0]["subscription_id"] + task = manager._subscriptions["conn-1"][sub_id].task + + await manager.handle_unsubscribe("conn-1", websocket, sub_id) + await asyncio.sleep(0) + + assert task.cancelled() or task.done() + assert any(f["type"] == "unsubscribed" for f in websocket.sent) + + async def test_remove_connection_cancels_every_task(self): + manager, _, _ = build_manager() + websocket = RecordingWebSocket() + + for sub_type in sorted(SUBSCRIPTION_TYPES): + await manager.handle_subscribe("conn-1", websocket, SUBSCRIBE_MESSAGES[sub_type]) + tasks = [s.task for s in manager._subscriptions["conn-1"].values()] + assert len(tasks) == len(SUBSCRIPTION_TYPES) + + manager.remove_connection("conn-1") + await asyncio.gather(*tasks, return_exceptions=True) + + assert all(t.done() for t in tasks) + assert "conn-1" not in manager._subscriptions + + async def test_shutdown_cancels_every_connection(self): + manager, _, _ = build_manager() + websocket = RecordingWebSocket() + + await manager.handle_subscribe("conn-1", websocket, SUBSCRIBE_MESSAGES["executors"]) + await manager.handle_subscribe("conn-2", websocket, SUBSCRIBE_MESSAGES["positions"]) + tasks = [ + s.task + for subs in manager._subscriptions.values() + for s in subs.values() + ] + + await manager.shutdown() + await asyncio.gather(*tasks, return_exceptions=True) + + assert all(t.done() for t in tasks) + assert manager._subscriptions == {} + + async def test_resubscribing_replaces_the_previous_task(self): + manager, _, _ = build_manager() + websocket = RecordingWebSocket() + + await manager.handle_subscribe("conn-1", websocket, SUBSCRIBE_MESSAGES["executors"]) + first = next(iter(manager._subscriptions["conn-1"].values())).task + await manager.handle_subscribe("conn-1", websocket, SUBSCRIBE_MESSAGES["executors"]) + await asyncio.sleep(0) + subs = manager._subscriptions["conn-1"] + assert len(subs) == 1 + second = next(iter(subs.values())).task + assert first is not second + assert first.done() or first.cancelled() + + manager.remove_connection("conn-1") + await asyncio.gather(first, second, return_exceptions=True) + + +class TestPerSubscriptionInterval: + """Each loop still sleeps its own subscription's interval.""" + + @pytest.fixture + def record_sleeps(self, monkeypatch): + real_sleep = asyncio.sleep + calls = [] + + async def fake_sleep(delay, *args, **kwargs): + calls.append(delay) + await real_sleep(0) + + monkeypatch.setattr(asyncio, "sleep", fake_sleep) + return calls + + async def test_two_subscriptions_sleep_their_own_intervals(self, record_sleeps): + manager, _, _ = build_manager() + fast_ws, slow_ws = RecordingWebSocket(), RecordingWebSocket() + fast = make_sub("executors", interval=0.5) + slow = make_sub("executor_summary", interval=7.0) + + tasks = [ + asyncio.create_task(manager._get_push_fn("executors")("c", fast_ws, fast)), + asyncio.create_task(manager._get_push_fn("executor_summary")("c", slow_ws, slow)), + ] + await fast_ws.wait_for_frames(1) + await slow_ws.wait_for_frames(1) + for task in tasks: + task.cancel() + await asyncio.gather(*tasks, return_exceptions=True) + + assert 0.5 in record_sleeps + assert 7.0 in record_sleeps + assert set(record_sleeps) <= {0.5, 7.0, 0} + + async def test_the_logs_loop_uses_its_own_interval_too(self, record_sleeps): + manager, _, _ = build_manager() + websocket = RecordingWebSocket() + sub = make_sub("executor_logs", interval=3.0, executor_id="e-1") + + task = asyncio.create_task( + manager._get_push_fn("executor_logs")("c", websocket, sub) + ) + await websocket.wait_for_frames(1) + task.cancel() + await asyncio.gather(task, return_exceptions=True) + + assert 3.0 in record_sleeps diff --git a/test/test_executor_ws_update_interval.py b/test/test_executor_ws_update_interval.py new file mode 100644 index 00000000..abe702eb --- /dev/null +++ b/test/test_executor_ws_update_interval.py @@ -0,0 +1,185 @@ +""" +Tests for the /ws/executors update-interval bounds (READ-057). + +The bounds used to be module-level literals in services/executor_ws_manager.py; +they now come from MarketDataSettings, so the MARKET_DATA_WS_EXECUTOR_* env vars +actually reach the executor WebSocket. + +Run with: pytest test/test_executor_ws_update_interval.py -v +""" +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from config import MarketDataSettings, settings +from services.executor_ws_manager import ExecutorWebSocketManager, _clamp_interval + + +@pytest.fixture +def custom_bounds(monkeypatch): + """Rebuild MarketDataSettings from the env, as an operator override would.""" + + def _apply(**env): + for key, value in env.items(): + monkeypatch.setenv(key, str(value)) + monkeypatch.setattr(settings, "market_data", MarketDataSettings()) + + return _apply + + +class TestClampIntervalDefaults: + """Default behaviour must be identical to the old hardcoded literals.""" + + def test_missing_interval_falls_back_to_two_seconds(self): + assert _clamp_interval(None) == 2.0 + + def test_interval_below_the_floor_is_raised_to_half_a_second(self): + assert _clamp_interval(0.01) == 0.5 + + def test_interval_above_the_ceiling_is_capped_at_sixty_seconds(self): + assert _clamp_interval(1000.0) == 60.0 + + def test_interval_inside_the_range_is_left_alone(self): + assert _clamp_interval(5.0) == 5.0 + + def test_executor_floor_is_stricter_than_the_market_data_floor(self): + md = MarketDataSettings() + assert md.ws_executor_min_update_interval > md.ws_min_update_interval + + +class TestClampIntervalIsConfigurable: + """The MARKET_DATA_WS_EXECUTOR_* env vars must change the clamping.""" + + def test_configured_floor_is_honoured(self, custom_bounds): + custom_bounds(MARKET_DATA_WS_EXECUTOR_MIN_UPDATE_INTERVAL=5.0) + assert _clamp_interval(1.0) == 5.0 + + def test_configured_ceiling_is_honoured(self, custom_bounds): + custom_bounds(MARKET_DATA_WS_EXECUTOR_MAX_UPDATE_INTERVAL=10.0) + assert _clamp_interval(30.0) == 10.0 + + def test_configured_default_is_honoured(self, custom_bounds): + custom_bounds(MARKET_DATA_WS_EXECUTOR_DEFAULT_UPDATE_INTERVAL=7.5) + assert _clamp_interval(None) == 7.5 + + def test_market_data_bounds_do_not_leak_into_the_executor_clamp(self, custom_bounds): + custom_bounds( + MARKET_DATA_WS_MIN_UPDATE_INTERVAL=0.01, + MARKET_DATA_WS_MAX_UPDATE_INTERVAL=3600.0, + ) + assert _clamp_interval(0.01) == 0.5 + assert _clamp_interval(3600.0) == 60.0 + + +class TestClampIntervalRejectsNonNumbers: + """A non-numeric interval must raise ValueError, not TypeError (CORR-113).""" + + @pytest.mark.parametrize( + "bad", + ["fast", "2", [1.0], {"seconds": 1.0}, (1.0,), object()], + ids=["str", "numeric_str", "list", "dict", "tuple", "object"], + ) + def test_non_numeric_interval_is_rejected(self, bad): + with pytest.raises(ValueError, match="update_interval must be a number"): + _clamp_interval(bad) + + @pytest.mark.parametrize("bad", [True, False], ids=["true", "false"]) + def test_booleans_are_not_accepted_as_numbers(self, bad): + with pytest.raises(ValueError, match="update_interval must be a number"): + _clamp_interval(bad) + + def test_the_error_names_the_offending_type(self): + with pytest.raises(ValueError, match="got str"): + _clamp_interval("fast") + + def test_integers_are_still_valid(self): + assert _clamp_interval(5) == 5.0 + + +class TestSubscribeAcknowledgement: + """handle_subscribe must apply and report the configured clamp.""" + + @staticmethod + def _manager(): + executor_service = MagicMock() + executor_service.get_executors = AsyncMock(return_value=[]) + return ExecutorWebSocketManager( + executor_service=executor_service, + market_data_service=MagicMock(), + ) + + @staticmethod + async def _subscribe(manager, websocket, msg): + await manager.handle_subscribe("conn-1", websocket, msg) + manager.remove_connection("conn-1") + + async def test_ack_reports_the_configured_floor(self, custom_bounds): + custom_bounds(MARKET_DATA_WS_EXECUTOR_MIN_UPDATE_INTERVAL=4.0) + manager = self._manager() + websocket = MagicMock() + websocket.send_json = AsyncMock() + + await self._subscribe(manager, websocket, {"type": "executors", "update_interval": 0.5}) + + ack = next( + call.args[0] + for call in websocket.send_json.call_args_list + if call.args[0].get("type") == "subscribed" + ) + assert ack["update_interval"] == 4.0 + + async def test_ack_reports_the_default_when_none_is_requested(self): + manager = self._manager() + websocket = MagicMock() + websocket.send_json = AsyncMock() + + await self._subscribe(manager, websocket, {"type": "executors"}) + + ack = next( + call.args[0] + for call in websocket.send_json.call_args_list + if call.args[0].get("type") == "subscribed" + ) + assert ack["update_interval"] == 2.0 + + +class TestSubscribeRejectsMalformedInterval: + """A malformed interval must come back as an error frame, not an exception.""" + + @staticmethod + def _manager(): + executor_service = MagicMock() + executor_service.get_executors = AsyncMock(return_value=[]) + return ExecutorWebSocketManager( + executor_service=executor_service, + market_data_service=MagicMock(), + ) + + @pytest.mark.parametrize( + "bad", + ["fast", ["1"], {"seconds": 1.0}, True], + ids=["str", "list", "dict", "bool"], + ) + async def test_malformed_interval_sends_an_error_frame(self, bad): + manager = self._manager() + websocket = MagicMock() + websocket.send_json = AsyncMock() + + await manager.handle_subscribe( + "conn-1", websocket, {"type": "executors", "update_interval": bad} + ) + + frames = [call.args[0] for call in websocket.send_json.call_args_list] + assert [f["type"] for f in frames] == ["error"] + assert "update_interval" in frames[0]["message"] + + async def test_malformed_interval_does_not_subscribe(self): + manager = self._manager() + websocket = MagicMock() + websocket.send_json = AsyncMock() + + await manager.handle_subscribe( + "conn-1", websocket, {"type": "executors", "update_interval": "fast"} + ) + + assert manager._subscriptions.get("conn-1", {}) == {} diff --git a/test/test_gas_token_backfill.py b/test/test_gas_token_backfill.py new file mode 100644 index 00000000..8add5718 --- /dev/null +++ b/test/test_gas_token_backfill.py @@ -0,0 +1,191 @@ +"""Historical liquidity rows get the gas token their chain actually pays in (CORR-104). + +Before ARCH-054 unified the chain -> native-gas-token map, two write paths damaged +gateway_clmm_events.gas_token and neither will revisit the rows: + + - the add/remove handlers' two-branch ternary ("SOL" if solana else "ETH" if ethereum + else None) wrote NULL on base/arbitrum/polygon, and a row inserted CONFIRMED is never + re-polled, so the NULL is permanent; + - the transaction poller's old 6-entry dict wrote the literal "UNKNOWN" for those same + chains. + +The parser is fixed going forward; these tests pin the one-shot repair of what it left +behind — including that it resolves through get_native_gas_token rather than a second copy +of the map, that it refuses to invent a token for a chain the map does not know, and that +running it twice is the same as running it once. +""" +from contextlib import asynccontextmanager +from decimal import Decimal + +import pytest +from sqlalchemy import create_engine +from sqlalchemy.orm import Session +from sqlalchemy.pool import StaticPool + +from database.models import GatewayCLMMEvent, GatewayCLMMPosition +from database.repositories.gateway_clmm_repository import GatewayCLMMRepository +from services.gateway_client import get_native_gas_token + + +class _AsyncSessionAdapter: + """The async surface the repository uses, over a real synchronous Session. + + aiosqlite is not installed in this environment, and the backfill only awaits + execute/flush — so this runs the real SELECT, the real join and the real UPDATE + against the real schema. Mocking the session away would prove nothing about a + query whose whole job is finding the right rows. + """ + + def __init__(self, session: Session): + self._session = session + + async def execute(self, statement): + return self._session.execute(statement) + + async def flush(self): + self._session.flush() + + +@pytest.fixture +def db(): + engine = create_engine("sqlite://", connect_args={"check_same_thread": False}, + poolclass=StaticPool) + GatewayCLMMPosition.__table__.create(engine) + GatewayCLMMEvent.__table__.create(engine) + session = Session(engine) + + def position(network): + row = GatewayCLMMPosition( + position_address=f"pos-{network}-{id(network)}", + pool_address="pool", + network=network, + connector="raydium/clmm", + wallet_address="wallet", + trading_pair="SOL-USDC", + base_token="SOL", + quote_token="USDC", + lower_price=Decimal("1"), + upper_price=Decimal("2"), + ) + session.add(row) + session.flush() + return row + + def event(network, gas_token, event_type="ADD_LIQUIDITY", gas_fee=Decimal("0.005")): + row = GatewayCLMMEvent( + position_id=position(network).id, + transaction_hash=f"tx-{network}-{gas_token}-{event_type}-{gas_fee}", + event_type=event_type, + gas_fee=gas_fee, + gas_token=gas_token, + status="CONFIRMED", + ) + session.add(row) + session.flush() + return row + + @asynccontextmanager + async def run(): + """One backfill run against the same database.""" + yield GatewayCLMMRepository(_AsyncSessionAdapter(session)) + + try: + yield type("DB", (), { + "event": staticmethod(event), + "run": staticmethod(run), + "session": session, + }) + finally: + session.close() + engine.dispose() + + +async def _backfill(db): + async with db.run() as repo: + return await repo.backfill_liquidity_gas_tokens() + + +@pytest.mark.asyncio +async def test_a_null_gas_token_is_filled_from_the_positions_chain(db): + """The add/remove ternary's damage: CONFIRMED with a fee, no currency named.""" + row = db.event("base-mainnet", None) + + report = await _backfill(db) + + assert row.gas_token == "ETH" + assert report["fixed"] == 1 + + +@pytest.mark.asyncio +async def test_an_unknown_gas_token_is_replaced(db): + """The poller's 6-entry dict answered "UNKNOWN" for base, arbitrum and polygon.""" + rows = [db.event("arbitrum-mainnet", "UNKNOWN"), db.event("polygon-mainnet", "UNKNOWN")] + + report = await _backfill(db) + + assert [row.gas_token for row in rows] == ["ETH", "MATIC"] + assert report["fixed"] == 2 + + +@pytest.mark.asyncio +async def test_an_already_correct_row_is_left_alone(db): + row = db.event("solana-mainnet-beta", "SOL") + + report = await _backfill(db) + + assert row.gas_token == "SOL" + assert report == {"fixed": 0, "unresolved": 0, "unresolved_networks": []} + + +@pytest.mark.asyncio +async def test_every_value_written_is_the_one_the_helper_returns(db): + """Acceptance criterion: the map is not reimplemented here.""" + rows = {chain: db.event(f"{chain}-mainnet", None) + for chain in ("solana", "ethereum", "polygon", "avalanche", + "optimism", "arbitrum", "base", "bsc", "cronos")} + + await _backfill(db) + + assert {chain: row.gas_token for chain, row in rows.items()} == { + chain: get_native_gas_token(chain) for chain in rows + } + + +@pytest.mark.asyncio +async def test_a_chain_the_map_does_not_know_is_reported_not_papered_over(db): + """An unmapped chain must surface as a count, not as a fabricated token.""" + row = db.event("hyperliquid-mainnet", None) + + report = await _backfill(db) + + assert row.gas_token is None + assert report["fixed"] == 0 + assert report["unresolved"] == 1 + assert report["unresolved_networks"] == ["hyperliquid-mainnet"] + + +@pytest.mark.asyncio +async def test_a_row_that_recorded_no_gas_fee_names_no_currency(db): + """A pending/feeless row's NULL is not damage: the routes write gas_token only + alongside a fee, and the poller still owns that row.""" + row = db.event("base-mainnet", None, gas_fee=None) + + report = await _backfill(db) + + assert row.gas_token is None + assert report["fixed"] == 0 + + +@pytest.mark.asyncio +async def test_running_it_twice_produces_the_same_result(db): + """Acceptance criterion: idempotent.""" + rows = [db.event("base-mainnet", None), db.event("arbitrum-mainnet", "UNKNOWN"), + db.event("solana-mainnet-beta", "SOL"), db.event("hyperliquid-mainnet", None)] + + first = await _backfill(db) + after_first = [row.gas_token for row in rows] + second = await _backfill(db) + + assert after_first == [row.gas_token for row in rows] == ["ETH", "ETH", "SOL", None] + assert first["fixed"] == 2 and second["fixed"] == 0 + assert first["unresolved"] == second["unresolved"] == 1 diff --git a/test/test_gateway_error_masking.py b/test/test_gateway_error_masking.py index 15f1259d..fe1c20cb 100644 --- a/test/test_gateway_error_masking.py +++ b/test/test_gateway_error_masking.py @@ -5,7 +5,7 @@ These tests pin the three places where treating that shape as data caused real damage (found in the 2026-07-13 audit): - /gateway/swap/quote rendered a Gateway 404 as a 200 quote with price "0", -- _refresh_position_data marked positions CLOSED in the DB on any Gateway error, +- refresh_position_data marked positions CLOSED in the DB on any Gateway error, - the transaction poller recorded transient Gateway errors as on-chain FAILED. Run with: pytest test/test_gateway_error_masking.py -v @@ -80,7 +80,7 @@ def test_swap_quote_gateway_unreachable_is_503(swap_app): # ============================================ -# _refresh_position_data must not close positions on Gateway errors +# refresh_position_data must not close positions on Gateway errors # ============================================ def _position(): @@ -95,7 +95,7 @@ def _position(): @pytest.mark.asyncio async def test_refresh_does_not_close_position_on_gateway_error(): - from routers.gateway_clmm import _refresh_position_data + from services.gateway_clmm_service import refresh_position_data accounts_service = _mock_accounts_service(clmm_positions_owned=ERROR_DICT) clmm_repo = SimpleNamespace( @@ -104,7 +104,7 @@ async def test_refresh_does_not_close_position_on_gateway_error(): update_position_fees=AsyncMock(), ) - await _refresh_position_data(_position(), accounts_service, clmm_repo) + await refresh_position_data(_position(), accounts_service.gateway_client, clmm_repo) clmm_repo.close_position.assert_not_awaited() @@ -113,7 +113,7 @@ async def test_refresh_skips_position_missing_from_valid_list(): """Absence from ONE positions-owned read is not proof of closure (RPC lag): the refresh skips the update and leaves close detection to the poller's consecutive-miss gate — a single read must never close a live position.""" - from routers.gateway_clmm import _refresh_position_data + from services.gateway_clmm_service import refresh_position_data accounts_service = _mock_accounts_service(clmm_positions_owned=[{"address": "OTHER"}]) clmm_repo = SimpleNamespace( @@ -122,7 +122,7 @@ async def test_refresh_skips_position_missing_from_valid_list(): update_position_fees=AsyncMock(), ) - await _refresh_position_data(_position(), accounts_service, clmm_repo) + await refresh_position_data(_position(), accounts_service.gateway_client, clmm_repo) clmm_repo.close_position.assert_not_awaited() clmm_repo.update_position_liquidity.assert_not_awaited() diff --git a/test/test_gateway_ping_cache.py b/test/test_gateway_ping_cache.py new file mode 100644 index 00000000..0268ba57 --- /dev/null +++ b/test/test_gateway_ping_cache.py @@ -0,0 +1,239 @@ +""" +Regression tests for the Gateway availability guard's ping cache (PERF-114). + +``deps.require_gateway_online`` guards ~39 routes and runs ahead of every one of them, so an +uncached ``GatewayClient.ping()`` charges each guarded request a full HTTP round-trip to Gateway. +These pin the cache that removes that cost, and — just as importantly — the ways it must NOT +lie: +- a burst of guarded requests inside the TTL costs one ping, not one per request, +- the verdict expires: a Gateway that goes down is reported unavailable within the TTL, +- any Gateway call that fails to connect drops the cached verdict immediately, so an outage + mid-TTL is never masked by a stale "available", +- the resulting 503 still carries the single-sourced detail, now with ``Retry-After``. + +Run with: pytest test/test_gateway_ping_cache.py -v +""" +import math +from types import SimpleNamespace + +import aiohttp +import pytest +from fastapi import Depends, FastAPI +from fastapi.testclient import TestClient + +import services.gateway_client as gateway_client_module +from deps import GATEWAY_UNAVAILABLE_DETAIL, get_accounts_service, require_gateway_online +from services.gateway_client import GatewayClient + +TTL = GatewayClient.PING_CACHE_TTL_SECONDS + + +class FakeClock: + """Stand-in for ``time.monotonic`` so TTL expiry is exercised without sleeping.""" + + def __init__(self): + self.now = 1000.0 + + def monotonic(self): + return self.now + + def advance(self, seconds): + self.now += seconds + + +@pytest.fixture +def clock(monkeypatch): + fake = FakeClock() + monkeypatch.setattr(gateway_client_module, "time", fake) + return fake + + +@pytest.fixture +def client_and_pings(monkeypatch): + """A GatewayClient whose root request is counted instead of hitting the network.""" + client = GatewayClient() + pings = [] + responses = {"status": "ok"} + + async def fake_request(method, path, params=None, json=None): + pings.append((method, path)) + return responses + + monkeypatch.setattr(client, "_request", fake_request) + return client, pings, responses + + +# ================================== +# The cache hit path (the whole point) +# ================================== + + +@pytest.mark.asyncio +async def test_repeated_pings_inside_the_ttl_cost_one_round_trip(client_and_pings, clock): + client, pings, _ = client_and_pings + + results = [await client.ping() for _ in range(10)] + + assert results == [True] * 10 + assert pings == [("GET", "")], "the guard must reuse the cached verdict, not re-ping per request" + + +@pytest.mark.asyncio +async def test_a_false_verdict_is_cached_too(client_and_pings, clock): + client, pings, responses = client_and_pings + responses["status"] = "down" + + assert await client.ping() is False + assert await client.ping() is False + assert len(pings) == 1 + + +# ============================================== +# The expiry path (a stale verdict must not stick) +# ============================================== + + +@pytest.mark.asyncio +async def test_the_verdict_expires_after_the_ttl(client_and_pings, clock): + client, pings, _ = client_and_pings + + assert await client.ping() is True + clock.advance(TTL + 0.01) + assert await client.ping() is True + + assert len(pings) == 2, "the verdict must not outlive PING_CACHE_TTL_SECONDS" + + +@pytest.mark.asyncio +async def test_a_gateway_that_goes_down_is_reported_within_the_ttl(client_and_pings, clock): + """An 'available' verdict may not survive the TTL once Gateway stops answering.""" + client, _, responses = client_and_pings + + assert await client.ping() is True + responses["status"] = "down" # Gateway dies right after the cached ping + + clock.advance(TTL) + assert await client.ping() is False + + +def test_the_ttl_stays_short_enough_to_be_a_guard(): + assert 0 < TTL <= 5, "the guard reports outages within the TTL; keep it in the seconds range" + + +# =============================================================== +# Invalidation: a failed Gateway call must not leave a stale 'up' +# =============================================================== + + +class _OkResponse: + """Minimal stand-in for the aiohttp response a healthy Gateway returns.""" + + ok = True + status = 200 + + async def __aenter__(self): + return self + + async def __aexit__(self, *exc_info): + return False + + async def json(self): + return {"status": "ok"} + + +class _FlakySession: + """Session that answers until ``state['down']`` flips, then refuses to connect.""" + + def __init__(self, state): + self.state = state + + def get(self, *args, **kwargs): + if self.state["down"]: + raise aiohttp.ClientConnectionError("connection refused") + return _OkResponse() + + post = get + delete = get + + +@pytest.mark.asyncio +async def test_a_failed_gateway_call_invalidates_the_cached_verdict(monkeypatch, clock): + """A connection failure on any call drops the cache, so the outage is not masked mid-TTL.""" + client = GatewayClient() + state = {"down": False} + + async def session(): + return _FlakySession(state) + + monkeypatch.setattr(client, "_get_session", session) + + assert await client.ping() is True + assert client._ping_cache is not None + + # Gateway goes down mid-TTL; an unrelated guarded handler's call is the one that finds out. + state["down"] = True + assert await client.get_wallets() is None + assert client._ping_cache is None, "an unreachable Gateway must clear the cached verdict" + + # ... and the very next guard (still inside the TTL) sees the outage, not the stale 'up'. + assert await client.ping() is False + + +@pytest.mark.asyncio +async def test_missing_mtls_certs_also_invalidate_the_cached_verdict(monkeypatch, clock): + client = GatewayClient() + client._ping_cache = (clock.monotonic() + TTL, True) + + async def no_certs(): + raise FileNotFoundError("gateway client cert missing") + + monkeypatch.setattr(client, "_get_session", no_certs) + + result = await client.get_wallets() + + assert result["status"] == 503 + assert client._ping_cache is None + + +# ========================================== +# The guard itself: ping count and the 503 +# ========================================== + + +@pytest.fixture +def guarded_client_for(): + app = FastAPI() + + @app.get("/guarded", dependencies=[Depends(require_gateway_online)]) + async def guarded(): + return {"ok": True} + + def build(gateway): + accounts_service = SimpleNamespace(gateway_client=gateway) + app.dependency_overrides[get_accounts_service] = lambda: accounts_service + return TestClient(app, raise_server_exceptions=False) + + return build + + +def test_many_guarded_requests_produce_one_ping(guarded_client_for, client_and_pings, clock): + gateway, pings, _ = client_and_pings + client = guarded_client_for(gateway) + + for _ in range(5): + assert client.get("/guarded").status_code == 200 + + assert len(pings) == 1, f"5 guarded requests inside the TTL pinged Gateway {len(pings)} times" + + +def test_the_503_carries_retry_after_and_the_shared_detail(guarded_client_for, client_and_pings, clock): + gateway, _, responses = client_and_pings + responses["status"] = "down" + client = guarded_client_for(gateway) + + response = client.get("/guarded") + + assert response.status_code == 503 + assert response.json()["detail"] == GATEWAY_UNAVAILABLE_DETAIL + assert response.headers["Retry-After"] == str(math.ceil(TTL)) + assert int(response.headers["Retry-After"]) >= 1 diff --git a/test/test_gateway_response_parsing.py b/test/test_gateway_response_parsing.py new file mode 100644 index 00000000..c2a6b565 --- /dev/null +++ b/test/test_gateway_response_parsing.py @@ -0,0 +1,336 @@ +"""Gateway's write-response wire format is parsed in one place, not copied per handler. + +Gateway answers a write with the transaction id under one of three keys (`signature` on +Solana, `txHash` on EVM, `hash` on some older shapes) and says nothing at all about which +token paid for the gas. Both facts used to be re-derived in every handler that recorded an +event, and the copies drifted: + +- the chain -> gas-token mapping existed three times — the 9-entry dict, a 6-entry one in + the poller, and a two-branch ternary (`"SOL" if chain == "solana" else "ETH" if chain == + "ethereum" else None`) in the add/remove handlers. On base/arbitrum/polygon that ternary + wrote gas_token NULL, and a row inserted CONFIRMED is never re-polled, so the NULL was + permanent. Add/remove liquidity gas costs were unusable off solana/ethereum; +- the tx-id extraction existed six times, and the CLMM open handler's copy read + `signature` alone — answering every EVM open with "no transaction signature returned" + for a position that had just been opened. + +These tests pin both invariants: one definition of each parser, and the same gas token +persisted by every CLMM event writer for the same chain. +""" +import ast +import pathlib +from types import SimpleNamespace +from unittest.mock import AsyncMock + +import pytest +from fastapi import FastAPI +from fastapi.testclient import TestClient + +from routers.gateway_extras import get_transaction_hash_from_response +from services.gateway_client import get_native_gas_token + +REPO_ROOT = pathlib.Path(__file__).resolve().parent.parent + +# Every module that writes a gas_token column or parses a Gateway write response. +GAS_TOKEN_WRITERS = [ + REPO_ROOT / "routers" / "gateway_clmm.py", + REPO_ROOT / "routers" / "gateway_amm.py", + REPO_ROOT / "routers" / "gateway_swap.py", + REPO_ROOT / "services" / "gateway_transaction_poller.py", + REPO_ROOT / "services" / "executor_service.py", +] + + +# ============================================ +# One definition of each parser +# ============================================ + +def _definitions_of(name: str): + found = [] + for directory in ("routers", "services"): + for path in sorted((REPO_ROOT / directory).rglob("*.py")): + tree = ast.parse(path.read_text()) + for node in ast.walk(tree): + if isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)) and node.name == name: + found.append(path.relative_to(REPO_ROOT).as_posix()) + return found + + +@pytest.mark.parametrize("parser, home", [ + ("get_transaction_status_from_response", "routers/gateway_extras.py"), + ("get_transaction_hash_from_response", "routers/gateway_extras.py"), + ("get_native_gas_token", "services/gateway_client.py"), +]) +def test_each_response_parser_is_defined_exactly_once(parser, home): + assert _definitions_of(parser) == [home] + + +def test_only_gateway_client_holds_a_chain_to_gas_token_map(): + """A second dict keyed by chain name is how the poller's 6-entry copy came to answer + "UNKNOWN" for base/bsc/cronos while the routers answered correctly.""" + offenders = [] + for directory in ("routers", "services"): + for path in sorted((REPO_ROOT / directory).rglob("*.py")): + for node in ast.walk(ast.parse(path.read_text())): + if not isinstance(node, ast.Dict): + continue + keys = {k.value for k in node.keys if isinstance(k, ast.Constant) and isinstance(k.value, str)} + if {"solana", "ethereum"} <= keys: + offenders.append(path.relative_to(REPO_ROOT).as_posix()) + assert offenders == ["services/gateway_client.py"] + + +def _gas_token_values(tree): + """Every expression assigned to a gas_token variable or dict key.""" + for node in ast.walk(tree): + if isinstance(node, ast.Assign): + for target in node.targets: + if isinstance(target, ast.Name) and target.id == "gas_token": + yield node.value + elif isinstance(node, ast.Dict): + for key, value in zip(node.keys, node.values): + if isinstance(key, ast.Constant) and key.value == "gas_token": + yield value + + +@pytest.mark.parametrize("path", GAS_TOKEN_WRITERS, ids=lambda p: p.name) +def test_every_gas_token_write_resolves_through_the_one_helper(path): + """No literal and no chain ternary: the value is either the helper's call or a name + bound from it. The old `"SOL" if chain == "solana" else ... else None` fails here.""" + tree = ast.parse(path.read_text()) + for value in _gas_token_values(tree): + names = {n.id for n in ast.walk(value) if isinstance(n, ast.Name)} + calls_helper = "get_native_gas_token" in names + # A bare `gas_token` / `status_result.get("gas_token")` reference is a name bound + # from the helper earlier in the same function (asserted for that binding above). + reads_a_binding = isinstance(value, (ast.Name, ast.Call, ast.Attribute)) and not isinstance( + value, ast.Constant) + assert calls_helper or reads_a_binding, ( + f"{path.name}:{value.lineno} writes gas_token from a literal or a chain " + f"ternary instead of get_native_gas_token()" + ) + + +# ============================================ +# The gas-token map itself +# ============================================ + +@pytest.mark.parametrize("chain, token", [ + ("solana", "SOL"), + ("ethereum", "ETH"), + ("polygon", "MATIC"), + ("avalanche", "AVAX"), + ("optimism", "ETH"), + ("arbitrum", "ETH"), + ("base", "ETH"), + ("bsc", "BNB"), + ("cronos", "CRO"), +]) +def test_the_gas_token_map_covers_every_supported_chain(chain, token): + assert get_native_gas_token(chain) == token + # Gateway is not consistent about casing in network ids. + assert get_native_gas_token(chain.upper()) == token + + +def test_an_unknown_chain_is_named_not_nulled(): + """"UNKNOWN" over None deliberately: a NULL column cannot be told apart from a row + written before gas was recorded at all, so the backfill can't find it.""" + assert get_native_gas_token("hyperliquid") == "UNKNOWN" + assert get_native_gas_token("") == "UNKNOWN" + + +# ============================================ +# Transaction-id extraction +# ============================================ + +@pytest.mark.parametrize("response, expected", [ + ({"signature": "5Uq9x"}, "5Uq9x"), # Solana + ({"txHash": "0xabc"}, "0xabc"), # EVM — the shape open used to drop + ({"hash": "0xdef"}, "0xdef"), # older Gateway shapes + ({"signature": "5Uq9x", "txHash": "0xabc"}, "5Uq9x"), # signature wins + ({"status": 1}, None), # nothing to report + ({"signature": None, "txHash": "0xabc"}, "0xabc"), # null signature falls through +]) +def test_the_transaction_id_is_read_from_whichever_key_gateway_used(response, expected): + assert get_transaction_hash_from_response(response) == expected + + +# ============================================ +# ADD/REMOVE persist the same gas token as COLLECT_FEES — the data bug +# ============================================ + +class _FakeRepo: + """Captures create_event payloads; every other repository call is a no-op await.""" + + events = [] + + def __init__(self, session): + pass + + async def get_position_by_address(self, position_address): + return SimpleNamespace( + id=1, + position_address=position_address, + pool_address="0xPOOL", + wallet_address="0xWALLET", + ) + + async def create_event(self, event_data): + _FakeRepo.events.append(event_data) + + def __getattr__(self, name): + return AsyncMock() + + +class _FakeDbManager: + def get_session_context(self): + class _Ctx: + async def __aenter__(self): + return SimpleNamespace() + + async def __aexit__(self, *exc): + return False + + return _Ctx() + + +def _fake_clmm_service(): + """The real service the routes write through, backed by _FakeRepo.""" + from services.gateway_clmm_service import GatewayCLMMService + + service = GatewayCLMMService(db_manager=_FakeDbManager()) + service.repository_class = _FakeRepo + return service + + +# A base-mainnet (EVM) write: no `signature`, gas paid in ETH. Under the old code this +# recorded gas_token NULL for add/remove and "ETH" for collect-fees — same chain, same +# wallet, same block. +def _evm_write(**data): + return {"txHash": "0xdeadbeef", "status": 1, "data": {"fee": 0.00042, **data}} + + +CLMM_WRITES = [ + ( + "/gateway/clmm/add", + {"connector": "uniswap", "network": "base-mainnet", "position_address": "0xPOS", + "base_token_amount": 0.5, "quote_token_amount": 50}, + "clmm_add_liquidity", + _evm_write(baseTokenAmountAdded=0.5, quoteTokenAmountAdded=50), + "ADD_LIQUIDITY", + ), + ( + "/gateway/clmm/remove", + {"connector": "uniswap", "network": "base-mainnet", "position_address": "0xPOS", + "percentage_to_remove": 50}, + "clmm_remove_liquidity", + _evm_write(baseTokenAmountRemoved=0.25, quoteTokenAmountRemoved=25), + "REMOVE_LIQUIDITY", + ), + ( + "/gateway/clmm/collect-fees", + {"connector": "uniswap", "network": "base-mainnet", "position_address": "0xPOS"}, + "clmm_collect_fees", + _evm_write(baseFeeAmountCollected=0.01, quoteFeeAmountCollected=1), + "COLLECT_FEES", + ), +] + + +def _record_clmm_write(route, body, gateway_method, gateway_response): + from deps import get_accounts_service, get_gateway_clmm_service + from routers import gateway_clmm + + _FakeRepo.events = [] + + gateway_client = SimpleNamespace( + ping=AsyncMock(return_value=True), + parse_network_id=lambda network_id: tuple(network_id.split("-", 1)), + get_wallet_address_or_default=AsyncMock(return_value="0xWALLET"), + clmm_positions_owned=AsyncMock(return_value=[]), + clmm_pool_info=AsyncMock(return_value={"price": 2500.0}), + **{gateway_method: AsyncMock(return_value=gateway_response)}, + ) + + app = FastAPI() + app.include_router(gateway_clmm.router) + app.dependency_overrides[get_accounts_service] = lambda: SimpleNamespace(gateway_client=gateway_client) + app.dependency_overrides[get_gateway_clmm_service] = lambda: _fake_clmm_service() + + response = TestClient(app, raise_server_exceptions=False).post(route, json=body) + assert response.status_code == 200, response.text + return _FakeRepo.events + + +@pytest.mark.parametrize("route, body, gateway_method, gateway_response, event_type", CLMM_WRITES) +def test_every_clmm_event_records_the_chains_gas_token( + route, body, gateway_method, gateway_response, event_type +): + events = _record_clmm_write(route, body, gateway_method, gateway_response) + + recorded = [e for e in events if e["event_type"] == event_type] + assert recorded, f"{event_type} wrote no event" + for event in recorded: + assert event["gas_token"] == "ETH", ( + f"{event_type} on base recorded gas_token {event['gas_token']!r}; a CONFIRMED " + "row is never re-polled, so a wrong value here is permanent" + ) + assert event["gas_fee"] == pytest.approx(0.00042) + # The EVM id came from `txHash`; reading `signature` alone got None here. + assert event["transaction_hash"] == "0xdeadbeef" + + +def test_add_remove_and_collect_fees_agree_on_the_gas_token(): + """The parity the ternary broke: three writers, one chain, one answer.""" + tokens = {} + for route, body, gateway_method, gateway_response, event_type in CLMM_WRITES: + events = _record_clmm_write(route, body, gateway_method, gateway_response) + tokens[event_type] = next(e["gas_token"] for e in events if e["event_type"] == event_type) + + assert set(tokens.values()) == {get_native_gas_token("base")}, tokens + assert None not in tokens.values() + + +def test_an_evm_open_is_not_rejected_for_having_no_solana_signature(): + """The open handler read `signature` alone, so a uniswap open that Gateway confirmed + came back as a 500 saying no signature was returned — for a position that existed.""" + from deps import get_accounts_service, get_gateway_clmm_service + from routers import gateway_clmm + + _FakeRepo.events = [] + + gateway_client = SimpleNamespace( + ping=AsyncMock(return_value=True), + parse_network_id=lambda network_id: tuple(network_id.split("-", 1)), + get_wallet_address_or_default=AsyncMock(return_value="0xWALLET"), + clmm_pool_info=AsyncMock(return_value={"price": 2500.0}), + clmm_open_position=AsyncMock(return_value={ + "txHash": "0xopened", + "status": 1, + "data": {"positionAddress": "0xNEWPOS", "fee": 0.00042}, + }), + ) + + app = FastAPI() + app.include_router(gateway_clmm.router) + app.dependency_overrides[get_accounts_service] = lambda: SimpleNamespace(gateway_client=gateway_client) + app.dependency_overrides[get_gateway_clmm_service] = lambda: _fake_clmm_service() + + response = TestClient(app, raise_server_exceptions=False).post("/gateway/clmm/open", json={ + "connector": "uniswap", + "network": "base-mainnet", + "pool_address": "0xPOOL", + "lower_price": 2000, + "upper_price": 3000, + "base_token_amount": 0.5, + }) + + assert response.status_code == 200, response.text + assert response.json()["transaction_hash"] == "0xopened" + + +def test_the_swap_router_reads_the_evm_transaction_id_too(): + """gateway_swap's copy of the extraction is gone; the shared helper is what runs.""" + source = (REPO_ROOT / "routers" / "gateway_swap.py").read_text() + assert 'result.get("signature")' not in source + assert "get_transaction_hash_from_response(result)" in source diff --git a/test/test_gateway_router_persistence.py b/test/test_gateway_router_persistence.py new file mode 100644 index 00000000..8b556def --- /dev/null +++ b/test/test_gateway_router_persistence.py @@ -0,0 +1,878 @@ +"""The rows a CLMM/swap route writes must not change when the writer moves. + +Session ownership moved out of `routers/gateway_clmm.py` and `routers/gateway_swap.py` +and into `GatewayCLMMService` / `GatewaySwapService` (ARCH-052). Those routes are the +only writers of `gateway_clmm_positions`, `gateway_clmm_events` and `gateway_swaps` — +real trading history that PnL and the transaction poller both read back — so a +refactor there is only safe if the same rows still land, with the same values. + +These tests drive the real FastAPI routes with a fake repository behind the service +and pin every column each handler writes, plus the two policies the move had to +preserve: a persistence failure never fails the trade, and "no such row" is still a +404 rather than a defaulted 200. +""" +from contextlib import asynccontextmanager +from decimal import Decimal +from types import SimpleNamespace +from unittest.mock import AsyncMock + +import pytest +from fastapi import FastAPI +from fastapi.testclient import TestClient + +from deps import get_accounts_service, get_gateway_clmm_service, get_gateway_swap_service +from routers import gateway_clmm, gateway_swap +from services.gateway_clmm_service import GatewayCLMMService +from services.gateway_swap_service import GatewaySwapService + +WALLET = "82SggYRE2Vo4jN4a2pk3aQ4SET4ctafZJGbowmCqyHx5" +POOL = "2sf5NYcY4zUPXUSmG6f66mskb24t5F8S11pC1Nz5nQT3" +POSITION = "9xQeWvG816bUx9EPjHmaT23yvVM2ZWbrrpZb9PusVFin" +SIGNATURE = "5xLmQ5s5xZ9jTqk3Y8bNvW2pR7cH4dF6gJ1kM3nP9qS8tU4vX6yZ2aB5cD7eF9gH1jK3zM5nP7qR9sT" +SOL = "So11111111111111111111111111111111111111112" +USDC = "EPjFWdd5AufqSSqeM2qN1xzybapC8G4wEGGkZwyTDt1v" + + +# --------------------------------------------------------------------------- +# Fakes +# --------------------------------------------------------------------------- + +class RepoCalls(list): + """Every repository call a request made, in order.""" + + def payload(self, method): + """The single payload passed to ``method`` (fails if not called exactly once).""" + matches = [args for name, args in self if name == method] + assert len(matches) == 1, f"{method} called {len(matches)} times, expected 1" + return matches[0] + + def names(self): + return [name for name, _ in self] + + +def _repo_class(calls, position=None, failing=False): + """A repository that records what it was asked to write.""" + + class _Repo: + def __init__(self, session): + if failing: + raise RuntimeError("database is down") + + async def create_position(self, position_data): + calls.append(("create_position", position_data)) + return SimpleNamespace(id=7) + + async def create_event(self, event_data): + calls.append(("create_event", event_data)) + return SimpleNamespace(id=11) + + async def get_position_by_address(self, address): + calls.append(("get_position_by_address", address)) + return position + + async def add_to_position_amounts(self, **kwargs): + calls.append(("add_to_position_amounts", kwargs)) + + async def subtract_from_position_amounts(self, **kwargs): + calls.append(("subtract_from_position_amounts", kwargs)) + + async def update_position_fees(self, **kwargs): + calls.append(("update_position_fees", kwargs)) + + async def update_position_liquidity(self, **kwargs): + calls.append(("update_position_liquidity", kwargs)) + + async def close_position(self, address, **kwargs): + calls.append(("close_position", {"position_address": address, **kwargs})) + + async def create_swap(self, swap_data): + calls.append(("create_swap", swap_data)) + return SimpleNamespace(id=3) + + async def get_swap_by_tx_hash(self, transaction_hash): + calls.append(("get_swap_by_tx_hash", transaction_hash)) + return position + + def to_dict(self, swap): + return {"transaction_hash": swap.transaction_hash} + + return _Repo + + +def _db_manager(): + manager = SimpleNamespace() + + @asynccontextmanager + async def session_context(): + yield object() + + manager.get_session_context = session_context + return manager + + +def _service(service_class, calls, position=None, failing=False): + service = service_class(db_manager=_db_manager()) + service.repository_class = _repo_class(calls, position=position, failing=failing) + return service + + +def _accounts_service(**gateway_client_methods): + gateway_client = SimpleNamespace( + ping=AsyncMock(return_value=True), + parse_network_id=lambda network_id: tuple(network_id.split("-", 1)), + get_wallet_address_or_default=AsyncMock(return_value=WALLET), + **{name: AsyncMock(return_value=value) for name, value in gateway_client_methods.items()}, + ) + return SimpleNamespace(gateway_client=gateway_client) + + +def _client(accounts_service, clmm_service=None, swap_service=None): + app = FastAPI() + app.include_router(gateway_clmm.router) + app.include_router(gateway_swap.router) + app.dependency_overrides[get_accounts_service] = lambda: accounts_service + app.dependency_overrides[get_gateway_clmm_service] = lambda: clmm_service + app.dependency_overrides[get_gateway_swap_service] = lambda: swap_service + return TestClient(app, raise_server_exceptions=False) + + +def _stored_position(**overrides): + """A position row as the repository hands it back.""" + return SimpleNamespace(**{ + "id": 7, + "position_address": POSITION, + "pool_address": POOL, + "wallet_address": WALLET, + "base_fee_collected": Decimal("0.5"), + "quote_fee_collected": Decimal("2.5"), + "base_token_amount": Decimal("0.0099"), + "quote_token_amount": Decimal("1.98"), + **overrides, + }) + + +# --------------------------------------------------------------------------- +# CLMM open: a position row and its OPEN event +# --------------------------------------------------------------------------- + +OPEN_BODY = { + "connector": "meteora", + "network": "solana-mainnet-beta", + "pool_address": POOL, + "lower_price": 150, + "upper_price": 250, + "base_token_amount": 0.01, + "quote_token_amount": 2, +} + +OPEN_RESULT = { + "signature": SIGNATURE, + "status": 1, + "data": { + "positionAddress": POSITION, + "positionRent": 0.05788, + "baseTokenAmountAdded": 0.0099, + "quoteTokenAmountAdded": 1.98, + "fee": 0.000011772, + }, +} + +POOL_INFO = {"baseTokenAddress": SOL, "quoteTokenAddress": USDC, "price": 200.0} + + +def test_open_writes_the_same_position_row(): + calls = RepoCalls() + accounts_service = _accounts_service(clmm_pool_info=POOL_INFO, clmm_open_position=OPEN_RESULT) + client = _client(accounts_service, clmm_service=_service(GatewayCLMMService, calls)) + + response = client.post("/gateway/clmm/open", json=OPEN_BODY) + assert response.status_code == 200 + + assert calls.payload("create_position") == { + "position_address": POSITION, + "pool_address": POOL, + "network": "solana-mainnet-beta", + "connector": "meteora", + "wallet_address": WALLET, + "trading_pair": f"{SOL}-{USDC}", + "base_token": SOL, + "quote_token": USDC, + "status": "OPEN", + "lower_price": 150.0, + "upper_price": 250.0, + # (upper - lower) / lower, computed on the request's Decimals + "percentage": float((Decimal("250") - Decimal("150")) / Decimal("150")), + "entry_price": 200.0, + "current_price": 200.0, + # The on-chain amounts, never the requested 0.01 / 2. + "initial_base_token_amount": 0.0099, + "initial_quote_token_amount": 1.98, + "position_rent": 0.05788, + "base_token_amount": 0.0099, + "quote_token_amount": 1.98, + "in_range": "UNKNOWN", + # The columns this route did not use to write at all. It shares one row + # builder with the poller's discovery sweep (ARCH-103), so the key set no + # longer depends on which path recorded the position — what the open route + # cannot know is NULL, and what is genuinely zero at open time is zero. + "lower_bin_id": None, + "upper_bin_id": None, + "base_fee_pending": 0.0, + "quote_fee_pending": 0.0, + "base_fee_collected": 0.0, + "quote_fee_collected": 0.0, + } + + +def test_open_writes_the_same_event_row(): + calls = RepoCalls() + accounts_service = _accounts_service(clmm_pool_info=POOL_INFO, clmm_open_position=OPEN_RESULT) + client = _client(accounts_service, clmm_service=_service(GatewayCLMMService, calls)) + + client.post("/gateway/clmm/open", json=OPEN_BODY) + + assert calls.payload("create_event") == { + # Keyed to the row create_position just returned, not to the address. + "position_id": 7, + "transaction_hash": SIGNATURE, + "event_type": "OPEN", + "base_token_amount": 0.0099, + "quote_token_amount": 1.98, + "gas_fee": 0.000011772, + "gas_token": "SOL", + "status": "CONFIRMED", + } + + +def test_open_answers_the_caller_even_when_the_write_fails(): + # The position is open on-chain; a bookkeeping failure must not be reported as a + # failed open. This policy now lives in one place, so this is what pins it. + accounts_service = _accounts_service(clmm_pool_info=POOL_INFO, clmm_open_position=OPEN_RESULT) + client = _client( + accounts_service, + clmm_service=_service(GatewayCLMMService, RepoCalls(), failing=True), + ) + + response = client.post("/gateway/clmm/open", json=OPEN_BODY) + assert response.status_code == 200 + assert response.json()["position_address"] == POSITION + assert response.json()["status"] == "confirmed" + + +# --------------------------------------------------------------------------- +# CLMM add / remove: event row plus the position bookkeeping +# --------------------------------------------------------------------------- + +def test_add_liquidity_writes_its_event_and_books_the_capital(): + calls = RepoCalls() + accounts_service = _accounts_service( + clmm_pool_info={"price": 205.0}, + clmm_add_liquidity={ + "signature": SIGNATURE, + "status": 1, + "data": {"baseTokenAmountAdded": 0.005, "quoteTokenAmountAdded": 1.0, "fee": 0.000009}, + }, + ) + client = _client( + accounts_service, + clmm_service=_service(GatewayCLMMService, calls, position=_stored_position()), + ) + + response = client.post("/gateway/clmm/add", json={ + "connector": "meteora", + "network": "solana-mainnet-beta", + "position_address": POSITION, + "base_token_amount": 0.005, + "quote_token_amount": 1, + }) + assert response.status_code == 200 + + assert calls.payload("create_event") == { + "position_id": 7, + "transaction_hash": SIGNATURE, + "event_type": "ADD_LIQUIDITY", + "base_token_amount": 0.005, + "quote_token_amount": 1.0, + "gas_fee": 0.000009, + "gas_token": "SOL", + "status": "CONFIRMED", + } + # The pool price read for the re-weighting still reaches the booking call. + assert calls.payload("add_to_position_amounts") == { + "position_address": POSITION, + "base_delta": Decimal("0.005"), + "quote_delta": Decimal("1.0"), + "entry_price": Decimal("205.0"), + } + + +def test_add_liquidity_books_nothing_while_the_transaction_is_only_submitted(): + # A SUBMITTED event is booked by the poller's confirm path instead; booking here + # too would double-count it. + calls = RepoCalls() + accounts_service = _accounts_service( + clmm_pool_info={"price": 205.0}, + clmm_add_liquidity={"signature": SIGNATURE, "status": 0}, + ) + client = _client( + accounts_service, + clmm_service=_service(GatewayCLMMService, calls, position=_stored_position()), + ) + + client.post("/gateway/clmm/add", json={ + "connector": "meteora", + "network": "solana-mainnet-beta", + "position_address": POSITION, + "base_token_amount": 0.005, + }) + + assert calls.payload("create_event")["status"] == "SUBMITTED" + assert "add_to_position_amounts" not in calls.names() + + +def test_remove_liquidity_writes_its_event_and_unbooks_the_capital(): + calls = RepoCalls() + accounts_service = _accounts_service(clmm_remove_liquidity={ + "signature": SIGNATURE, + "status": 1, + "data": {"baseTokenAmountRemoved": 0.004, "quoteTokenAmountRemoved": 0.8, "fee": 0.000008}, + }) + client = _client( + accounts_service, + clmm_service=_service(GatewayCLMMService, calls, position=_stored_position()), + ) + + response = client.post("/gateway/clmm/remove", json={ + "connector": "meteora", + "network": "solana-mainnet-beta", + "position_address": POSITION, + "percentage_to_remove": 50, + }) + assert response.status_code == 200 + + # No "percentage" key: GatewayCLMMEvent has no such column. + assert calls.payload("create_event") == { + "position_id": 7, + "transaction_hash": SIGNATURE, + "event_type": "REMOVE_LIQUIDITY", + "base_token_amount": 0.004, + "quote_token_amount": 0.8, + "gas_fee": 0.000008, + "gas_token": "SOL", + "status": "CONFIRMED", + } + assert calls.payload("subtract_from_position_amounts") == { + "position_address": POSITION, + "base_delta": Decimal("0.004"), + "quote_delta": Decimal("0.8"), + } + + +def test_an_event_for_an_unknown_position_is_skipped_not_invented(): + calls = RepoCalls() + accounts_service = _accounts_service(clmm_remove_liquidity={"signature": SIGNATURE, "status": 1}) + client = _client( + accounts_service, + clmm_service=_service(GatewayCLMMService, calls, position=None), + ) + + response = client.post("/gateway/clmm/remove", json={ + "connector": "meteora", + "network": "solana-mainnet-beta", + "position_address": POSITION, + "percentage_to_remove": 50, + }) + + assert response.status_code == 200 + assert "create_event" not in calls.names() + + +# --------------------------------------------------------------------------- +# CLMM collect-fees and close: the fee accounting +# --------------------------------------------------------------------------- + +def test_collect_fees_writes_its_event_and_rolls_the_collected_totals(): + calls = RepoCalls() + accounts_service = _accounts_service( + clmm_positions_owned=[{"address": POSITION, "baseFeeAmount": 0.01, "quoteFeeAmount": 2.0}], + clmm_collect_fees={ + "signature": SIGNATURE, + "status": 1, + "data": {"baseFeeAmountCollected": 0.01, "quoteFeeAmountCollected": 2.0, "fee": 0.000005}, + }, + ) + client = _client( + accounts_service, + clmm_service=_service(GatewayCLMMService, calls, position=_stored_position()), + ) + + response = client.post("/gateway/clmm/collect-fees", json={ + "connector": "meteora", + "network": "solana-mainnet-beta", + "position_address": POSITION, + }) + assert response.status_code == 200 + + assert calls.payload("create_event") == { + "position_id": 7, + "transaction_hash": SIGNATURE, + "event_type": "COLLECT_FEES", + "base_fee_collected": 0.01, + "quote_fee_collected": 2.0, + "gas_fee": 0.000005, + "gas_token": "SOL", + "status": "CONFIRMED", + } + # Added to what the row already held (0.5 / 2.5), and pending reset to zero. + assert calls.payload("update_position_fees") == { + "position_address": POSITION, + "base_fee_collected": Decimal("0.51"), + "quote_fee_collected": Decimal("4.5"), + "base_fee_pending": Decimal("0"), + "quote_fee_pending": Decimal("0"), + } + + +def test_close_writes_its_event_and_closes_the_row_once_gateway_agrees(monkeypatch): + import services.gateway_clmm_service as service_module + + # The close path waits for the transaction to propagate before verifying. + monkeypatch.setattr(service_module.asyncio, "sleep", AsyncMock()) + + calls = RepoCalls() + accounts_service = _accounts_service( + clmm_positions_owned=[{"address": POSITION, "baseFeeAmount": 0.01, + "quoteFeeAmount": 2.0, "price": 198.0}], + clmm_close_position={ + "signature": SIGNATURE, + "status": 1, + "data": { + "baseTokenAmountRemoved": 0.0099, + "quoteTokenAmountRemoved": 1.98, + "baseFeeAmountCollected": 0.01, + "quoteFeeAmountCollected": 2.0, + "positionRentRefunded": 0.05788, + "fee": 0.000011, + }, + }, + # Gateway no longer knows the position: proof the close landed. + clmm_position_info={"error": "Position not found", "status": 404}, + ) + client = _client( + accounts_service, + clmm_service=_service(GatewayCLMMService, calls, position=_stored_position()), + ) + + response = client.post("/gateway/clmm/close", json={ + "connector": "meteora", + "network": "solana-mainnet-beta", + "position_address": POSITION, + }) + assert response.status_code == 200 + + assert calls.payload("create_event") == { + "position_id": 7, + "transaction_hash": SIGNATURE, + "event_type": "CLOSE", + "base_token_amount": 0.0099, + "quote_token_amount": 1.98, + "base_fee_collected": 0.01, + "quote_fee_collected": 2.0, + "gas_fee": 0.000011, + "gas_token": "SOL", + "status": "CONFIRMED", + } + assert calls.payload("update_position_liquidity") == { + "position_address": POSITION, + "base_token_amount": Decimal("0.0099"), + "quote_token_amount": Decimal("1.98"), + "current_price": Decimal("198.0"), + } + assert calls.payload("close_position") == { + "position_address": POSITION, + "position_rent_refunded": Decimal("0.05788"), + } + + +def test_close_releases_the_session_before_waiting_for_propagation(monkeypatch): + """The two-second propagation wait must not hold a pooled connection. + + The close path books the fees, lets the session go, waits, and only then opens a + second short session to mark the row CLOSED. Holding one connection idle per + close is how a fleet closing several positions at once drains the pool while the + database has nothing to do (PERF-105). + """ + import services.gateway_clmm_service as service_module + + depth = [] # sessions currently open + timeline = [] # what happened, in order + sessions_open_during_sleep = [] + + @asynccontextmanager + async def session_context(): + depth.append(1) + timeline.append("session_open") + try: + yield object() + finally: + depth.pop() + timeline.append("session_close") + + async def fake_sleep(_seconds): + timeline.append("sleep") + sessions_open_during_sleep.append(len(depth)) + + monkeypatch.setattr(service_module.asyncio, "sleep", fake_sleep) + + calls = RepoCalls() + accounts_service = _accounts_service( + clmm_positions_owned=[{"address": POSITION, "baseFeeAmount": 0.01, + "quoteFeeAmount": 2.0, "price": 198.0}], + clmm_close_position={ + "signature": SIGNATURE, + "status": 1, + "data": { + "baseTokenAmountRemoved": 0.0099, + "quoteTokenAmountRemoved": 1.98, + "baseFeeAmountCollected": 0.01, + "quoteFeeAmountCollected": 2.0, + "positionRentRefunded": 0.05788, + "fee": 0.000011, + }, + }, + clmm_position_info={"error": "Position not found", "status": 404}, + ) + service = GatewayCLMMService(db_manager=SimpleNamespace(get_session_context=session_context)) + service.repository_class = _repo_class(calls, position=_stored_position()) + client = _client(accounts_service, clmm_service=service) + + response = client.post("/gateway/clmm/close", json={ + "connector": "meteora", + "network": "solana-mainnet-beta", + "position_address": POSITION, + }) + assert response.status_code == 200 + + # No connection was checked out while we waited on the chain. + assert sessions_open_during_sleep == [0] + # The bookkeeping commits, then the wait, then a second session for the close. + # (Earlier pairs belong to the route's own wallet lookup, before the close.) + assert timeline.count("sleep") == 1 + assert timeline[-3:] == ["sleep", "session_open", "session_close"] + # And the same rows still land, in the same order. + assert calls.names()[-5:] == [ + "get_position_by_address", "create_event", + "update_position_fees", "update_position_liquidity", + "close_position", + ] + + +def test_a_failed_close_mutates_nothing_but_still_files_the_event(): + calls = RepoCalls() + accounts_service = _accounts_service( + clmm_positions_owned=[{"address": POSITION, "baseFeeAmount": 0.01, "quoteFeeAmount": 2.0}], + clmm_close_position={"signature": SIGNATURE, "status": -1, "data": {"fee": 0.000011}}, + ) + client = _client( + accounts_service, + clmm_service=_service(GatewayCLMMService, calls, position=_stored_position()), + ) + + response = client.post("/gateway/clmm/close", json={ + "connector": "meteora", + "network": "solana-mainnet-beta", + "position_address": POSITION, + }) + assert response.status_code == 200 + + assert calls.payload("create_event")["status"] == "FAILED" + # No fee booking, no price update, no close: a reverted close changes nothing. + assert "update_position_fees" not in calls.names() + assert "close_position" not in calls.names() + + +# --------------------------------------------------------------------------- +# Swaps +# --------------------------------------------------------------------------- + +def test_execute_swap_writes_the_same_swap_row(): + calls = RepoCalls() + accounts_service = _accounts_service(execute_swap={ + "signature": SIGNATURE, + "status": 1, + "data": { + "amountIn": 0.01, + "amountOut": 0.878444, + "fee": 0.000005, + "slippagePct": 0.5, + "poolAddress": POOL, + }, + }) + client = _client(accounts_service, swap_service=_service(GatewaySwapService, calls)) + + response = client.post("/gateway/swap/execute", json={ + "connector": "jupiter/router", + "network": "solana-mainnet-beta", + "trading_pair": "SOL-USDC", + "side": "SELL", + "amount": 0.01, + "slippage_pct": 1, + }) + assert response.status_code == 200 + + assert calls.payload("create_swap") == { + "transaction_hash": SIGNATURE, + "network": "solana-mainnet-beta", + # The base venue name: "jupiter/router" files under "jupiter". + "connector": "jupiter", + "wallet_address": WALLET, + "trading_pair": "SOL-USDC", + "base_token": "SOL", + "quote_token": "USDC", + "side": "SELL", + "input_amount": 0.01, + "output_amount": 0.878444, + "price": float(Decimal("0.878444") / Decimal("0.01")), + # What Gateway says it applied, not the 1 that was asked for. + "slippage_pct": 0.5, + "gas_fee": 0.000005, + "gas_token": "SOL", + "status": "CONFIRMED", + "pool_address": POOL, + } + + +def test_a_submitted_swap_records_placeholders_and_no_gas(): + calls = RepoCalls() + accounts_service = _accounts_service(execute_swap={"signature": SIGNATURE, "status": 0}) + client = _client(accounts_service, swap_service=_service(GatewaySwapService, calls)) + + response = client.post("/gateway/swap/execute", json={ + "connector": "jupiter", + "network": "solana-mainnet-beta", + "trading_pair": "SOL-USDC", + "side": "BUY", + "amount": 0.01, + }) + assert response.status_code == 200 + + row = calls.payload("create_swap") + assert row["status"] == "SUBMITTED" + # BUY: the requested amount is the base leg out; the unknown leg stays 0. + assert (row["input_amount"], row["output_amount"], row["price"]) == (0.0, 0.01, 0.0) + assert row["gas_fee"] is None and row["gas_token"] is None + assert row["pool_address"] is None + # The response says nothing about a fill it does not know. + assert response.json()["output_amount"] is None + + +def test_a_swap_is_still_reported_when_the_write_fails(): + accounts_service = _accounts_service(execute_swap={ + "signature": SIGNATURE, "status": 1, + "data": {"amountIn": 0.01, "amountOut": 0.878444}, + }) + client = _client( + accounts_service, + swap_service=_service(GatewaySwapService, RepoCalls(), failing=True), + ) + + response = client.post("/gateway/swap/execute", json={ + "connector": "jupiter", + "network": "solana-mainnet-beta", + "trading_pair": "SOL-USDC", + "side": "SELL", + "amount": 0.01, + }) + + assert response.status_code == 200 + assert response.json()["transaction_hash"] == SIGNATURE + + +def test_an_unknown_swap_is_a_404_and_an_unreachable_database_is_not(): + # The trap of routing reads through a helper that swallows exceptions and returns + # a default: a database outage would answer 404 "Swap not found", which reads as + # "that swap never happened". + accounts_service = _accounts_service() + + client = _client(accounts_service, swap_service=_service(GatewaySwapService, RepoCalls())) + assert client.get(f"/gateway/swaps/{SIGNATURE}/status").status_code == 404 + + client = _client( + accounts_service, + swap_service=_service(GatewaySwapService, RepoCalls(), failing=True), + ) + assert client.get(f"/gateway/swaps/{SIGNATURE}/status").status_code == 500 + + +def test_a_known_swap_is_returned_as_the_repository_renders_it(): + swap = SimpleNamespace(transaction_hash=SIGNATURE) + client = _client( + _accounts_service(), + swap_service=_service(GatewaySwapService, RepoCalls(), position=swap), + ) + + response = client.get(f"/gateway/swaps/{SIGNATURE}/status") + assert response.status_code == 200 + assert response.json() == {"transaction_hash": SIGNATURE} + + +# --------------------------------------------------------------------------- +# The failed-write path keeps its single transaction-id parser +# --------------------------------------------------------------------------- + +@pytest.mark.asyncio +async def test_a_reverted_write_is_filed_against_the_position_it_belonged_to(): + from services.gateway_client import GatewayError + + calls = RepoCalls() + service = _service(GatewayCLMMService, calls, position=_stored_position()) + error = GatewayError(f"Transaction {SIGNATURE} landed on-chain but failed: 0x1771", status=500) + + await gateway_clmm._record_failed_write( + service, error, event_type="CLOSE", position_address=POSITION) + + assert calls.payload("create_event") == { + "position_id": 7, + "transaction_hash": SIGNATURE, + "event_type": "CLOSE", + "status": "FAILED", + "error_message": str(error), + } + + +@pytest.mark.asyncio +async def test_a_failure_that_never_reached_the_chain_writes_nothing(): + from services.gateway_client import GatewayError + + calls = RepoCalls() + service = _service(GatewayCLMMService, calls, position=_stored_position()) + + await gateway_clmm._record_failed_write( + service, + GatewayError("Simulation failed: insufficient funds", status=400), + event_type="CLOSE", + position_address=POSITION, + ) + + assert calls.names() == [] + + +# --------------------------------------------------------------------------- +# The read envelopes the handlers used to hand-build +# --------------------------------------------------------------------------- + +def _search_repo_class(calls, positions, refreshed): + class _Repo: + def __init__(self, session): + pass + + async def get_positions(self, **kwargs): + calls.append(("get_positions", kwargs)) + return positions + + async def get_position_by_address(self, address): + return _stored_position(position_address=address, connector="meteora", + network="solana-mainnet-beta") + + async def get_position_events(self, **kwargs): + calls.append(("get_position_events", kwargs)) + return [SimpleNamespace(id=1), SimpleNamespace(id=2)] + + def position_to_dict(self, position): + return {"position_address": position.position_address} + + def event_to_dict(self, event): + return {"id": event.id} + + async def update_position_liquidity(self, **kwargs): + refreshed.append(kwargs) + + async def update_position_fees(self, **kwargs): + pass + + async def close_position(self, address, **kwargs): + pass + + return _Repo + + +def _search_service(calls, positions, refreshed=None): + service = GatewayCLMMService(db_manager=_db_manager()) + service.repository_class = _search_repo_class(calls, positions, refreshed if refreshed is not None else []) + return service + + +def test_position_search_answers_the_same_paginated_envelope(): + calls = RepoCalls() + positions = [_stored_position(position_address=f"POS-{i}") for i in range(2)] + client = _client(_accounts_service(), clmm_service=_search_service(calls, positions)) + + response = client.post("/gateway/clmm/positions/search?limit=2000&offset=10") + + assert response.status_code == 200 + assert response.json() == { + "data": [{"position_address": "POS-0"}, {"position_address": "POS-1"}], + "pagination": { + # Clamped to the 1000 ceiling, and echoed clamped. + "limit": 1000, + "offset": 10, + "has_more": False, + "total_count": 12, + }, + } + assert calls.payload("get_positions")["limit"] == 1000 + + +def test_a_full_page_reports_more_rather_than_a_wrong_total(): + calls = RepoCalls() + positions = [_stored_position(position_address=f"POS-{i}") for i in range(2)] + client = _client(_accounts_service(), clmm_service=_search_service(calls, positions)) + + body = client.post("/gateway/clmm/positions/search?limit=2").json() + + assert body["pagination"]["has_more"] is True + assert body["pagination"]["total_count"] is None + + +def test_a_refreshing_search_writes_back_what_gateway_reports(): + calls, refreshed = RepoCalls(), [] + positions = [_stored_position(position_address=POSITION, connector="meteora", + network="solana-mainnet-beta")] + accounts_service = _accounts_service(clmm_positions_owned=[{ + "address": POSITION, + "price": 200.0, + "lowerPrice": 150.0, + "upperPrice": 250.0, + "baseTokenAmount": 0.0099, + "quoteTokenAmount": 1.98, + }]) + client = _client( + accounts_service, + clmm_service=_search_service(calls, positions, refreshed), + ) + + response = client.post("/gateway/clmm/positions/search?refresh=true") + + assert response.status_code == 200 + assert refreshed == [{ + "position_address": POSITION, + "base_token_amount": Decimal("0.0099"), + "quote_token_amount": Decimal("1.98"), + "in_range": "IN_RANGE", + "current_price": Decimal("200.0"), + }] + + +def test_position_events_answer_the_same_envelope(): + calls = RepoCalls() + client = _client(_accounts_service(), clmm_service=_search_service(calls, [])) + + response = client.get(f"/gateway/clmm/positions/{POSITION}/events?event_type=CLOSE&limit=5") + + assert response.status_code == 200 + assert response.json() == {"data": [{"id": 1}, {"id": 2}], "total_count": 2} + assert calls.payload("get_position_events") == { + "position_address": POSITION, + "event_type": "CLOSE", + "limit": 5, + } diff --git a/test/test_market_data_ws_push_loops.py b/test/test_market_data_ws_push_loops.py new file mode 100644 index 00000000..42823bad --- /dev/null +++ b/test/test_market_data_ws_push_loops.py @@ -0,0 +1,264 @@ +""" +Tests for the /ws/market-data push loops (CORR-102). + +The candles, order-book and trades loops used to wrap their whole poll body in +`except (WebSocketDisconnect, RuntimeError)` and break, so a RuntimeError raised +by the *data fetch* (a service fault, a connector still initialising) killed the +subscription for good and logged a disconnect that never happened. All three now +go through one `_send_or_stop` helper that guards only `send_json`; fetch errors +fall to `except Exception`, which logs and retries on the next interval. + +Run with: pytest test/test_market_data_ws_push_loops.py -v --asyncio-mode=auto +""" +import asyncio +from unittest.mock import AsyncMock, MagicMock + +import pandas as pd +import pytest +from fastapi.websockets import WebSocketDisconnect + +from services.websocket_manager import Subscription, WebSocketManager + + +class RecordingWebSocket: + """Captures pushed frames; can be told to raise on the first N sends.""" + + def __init__(self, raise_on_send=None, fail_first=0): + self.sent = [] + self.raise_on_send = raise_on_send + self.fail_first = fail_first + self.send_calls = 0 + self._target = None + self._reached = asyncio.Event() + + async def send_json(self, message): + self.send_calls += 1 + if self.raise_on_send is not None and self.send_calls > self.fail_first: + raise self.raise_on_send + self.sent.append(message) + if self._target is not None and len(self.sent) >= self._target: + self._reached.set() + + async def wait_for_frames(self, count, timeout=2.0): + self._target = count + if len(self.sent) >= count: + return + self._reached.clear() + await asyncio.wait_for(self._reached.wait(), timeout) + + +class FlakyFetch: + """Raises RuntimeError on the first `failures` calls, then returns `value`.""" + + def __init__(self, value, failures=1): + self.value = value + self.failures = failures + self.calls = 0 + + def __call__(self, *args, **kwargs): + self.calls += 1 + if self.calls <= self.failures: + raise RuntimeError("market data service is not ready") + return self.value + + +def make_candles_df(timestamp=1_700_000_000.0): + return pd.DataFrame( + [{"timestamp": timestamp, "open": 1.0, "high": 2.0, "low": 0.5, "close": 1.5, "volume": 10.0}] + ) + + +def make_feed(df=None): + feed = MagicMock() + feed.ready = True + feed.candles_df = make_candles_df() if df is None else df + return feed + + +def make_order_book(): + ob = MagicMock() + ob.last_diff_uid = None + ob.snapshot_uid = None + bids = pd.DataFrame([{"price": 100.0, "amount": 1.0}]) + asks = pd.DataFrame([{"price": 101.0, "amount": 2.0}]) + ob.snapshot = (bids, asks) + return ob + + +def make_manager(): + market_data_service = MagicMock() + market_data_service.get_candles_feed = AsyncMock(return_value=make_feed()) + market_data_service.get_order_book = MagicMock(return_value=make_order_book()) + return WebSocketManager(market_data_service), market_data_service + + +def make_sub(sub_type, interval=0.01): + return Subscription( + subscription_id=f"{sub_type}_binance_BTC-USDT", + sub_type=sub_type, + connector="binance", + trading_pair="BTC-USDT", + update_interval=interval, + interval="1m", + max_records=100, + depth=10, + ) + + +# --------------------------------------------------------------------------- +# A RuntimeError from the fetch must NOT end the subscription +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_candles_loop_survives_a_runtime_error_from_the_fetch(): + manager, service = make_manager() + flaky = FlakyFetch(make_feed(), failures=2) + + async def get_candles_feed(config): + return flaky(config) + + service.get_candles_feed = AsyncMock(side_effect=get_candles_feed) + ws = RecordingWebSocket() + sub = make_sub("candles") + + task = asyncio.create_task(manager._candles_push_loop(ws, sub)) + await ws.wait_for_frames(1) + task.cancel() + + # The two failing fetches were logged and retried, not fatal. + assert flaky.calls >= 3 + assert ws.sent[0]["type"] == "candles" + assert not task.done() or task.cancelled() + + +@pytest.mark.asyncio +async def test_order_book_loop_survives_a_runtime_error_from_the_fetch(): + manager, service = make_manager() + flaky = FlakyFetch(make_order_book(), failures=2) + service.get_order_book = MagicMock(side_effect=flaky) + ws = RecordingWebSocket() + sub = make_sub("order_book") + + task = asyncio.create_task(manager._order_book_push_loop(ws, sub)) + await ws.wait_for_frames(1) + task.cancel() + + assert flaky.calls >= 3 + assert ws.sent[0]["type"] == "order_book" + assert ws.sent[0]["data"] == {"bids": [[100.0, 1.0]], "asks": [[101.0, 2.0]]} + + +class FlakyBuffer(list): + """A trade buffer whose first `failures` drains raise RuntimeError.""" + + def __init__(self, items, failures=2): + super().__init__(items) + self.failures = failures + self.drains = 0 + + def __getitem__(self, item): + if isinstance(item, slice): + self.drains += 1 + if self.drains <= self.failures: + raise RuntimeError("transient fault while draining") + return list.__getitem__(self, item) + + +@pytest.mark.asyncio +async def test_trades_loop_survives_a_runtime_error_while_draining(): + """A RuntimeError raised before the send is retried, not fatal.""" + manager, _ = make_manager() + ws = RecordingWebSocket() + sub = make_sub("trades") + sub.trade_buffer = FlakyBuffer([{"price": 1.0, "amount": 2.0}], failures=2) + + task = asyncio.create_task(manager._trades_push_loop(ws, sub)) + await ws.wait_for_frames(1) + task.cancel() + + assert sub.trade_buffer.drains >= 3 + assert ws.sent[0]["type"] == "trades" + assert ws.sent[0]["data"] == [{"price": 1.0, "amount": 2.0}] + + +# --------------------------------------------------------------------------- +# A dropped client still ends the loop, once +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +@pytest.mark.parametrize("error", [WebSocketDisconnect(), RuntimeError("client disconnected")]) +async def test_candles_loop_stops_when_the_client_is_gone(error): + manager, _ = make_manager() + ws = RecordingWebSocket(raise_on_send=error) + sub = make_sub("candles") + + await asyncio.wait_for(manager._candles_push_loop(ws, sub), timeout=1.0) + + assert ws.send_calls == 1 # logged once, loop ended + assert ws.sent == [] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("error", [WebSocketDisconnect(), RuntimeError("client disconnected")]) +async def test_order_book_loop_stops_when_the_client_is_gone(error): + manager, _ = make_manager() + ws = RecordingWebSocket(raise_on_send=error) + sub = make_sub("order_book") + + await asyncio.wait_for(manager._order_book_push_loop(ws, sub), timeout=1.0) + + assert ws.send_calls == 1 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("error", [WebSocketDisconnect(), RuntimeError("client disconnected")]) +async def test_trades_loop_stops_when_the_client_is_gone(error): + manager, _ = make_manager() + ws = RecordingWebSocket(raise_on_send=error) + sub = make_sub("trades") + sub.trade_buffer.append({"price": 1.0, "amount": 2.0}) + + await asyncio.wait_for(manager._trades_push_loop(ws, sub), timeout=1.0) + + assert ws.send_calls == 1 + + +# --------------------------------------------------------------------------- +# All three loops go through the one shared helper +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_send_or_stop_reports_true_on_a_healthy_send(): + manager, _ = make_manager() + ws = RecordingWebSocket() + sub = make_sub("candles") + + assert await manager._send_or_stop(ws, sub, "candles", {"type": "candles"}) is True + assert ws.sent == [{"type": "candles"}] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("error", [WebSocketDisconnect(), RuntimeError("boom")]) +async def test_send_or_stop_reports_false_when_the_send_fails(error): + manager, _ = make_manager() + ws = RecordingWebSocket(raise_on_send=error) + sub = make_sub("candles") + + assert await manager._send_or_stop(ws, sub, "candles", {"type": "candles"}) is False + + +def test_no_push_loop_guards_its_fetch_against_runtime_error(): + """The regression guard: the broad guard must not come back.""" + import inspect + + import services.websocket_manager as module + + source = inspect.getsource(module) + assert "except (WebSocketDisconnect, RuntimeError)" in source # still guarded in the helper + assert source.count("except (WebSocketDisconnect, RuntimeError)") == 1 + assert inspect.getsource(module.WebSocketManager._send_or_stop).count( + "except (WebSocketDisconnect, RuntimeError)" + ) == 1 diff --git a/test/test_mqtt_log_dedup.py b/test/test_mqtt_log_dedup.py new file mode 100644 index 00000000..9c690dbc --- /dev/null +++ b/test/test_mqtt_log_dedup.py @@ -0,0 +1,76 @@ +"""Tests for the bounded, ordered log-deduplication cache in MQTTManager.""" + +import time +from collections import OrderedDict + +from utils.mqtt_manager import MQTTManager + + +def make_manager() -> MQTTManager: + return MQTTManager(host="localhost", port=1883, username="u", password="p") + + +class CountingOrderedDict(OrderedDict): + """OrderedDict that records how many times the cleanup loop peeks at it.""" + + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + self.values_calls = 0 + + def values(self): + self.values_calls += 1 + return super().values() + + +async def test_duplicate_log_messages_within_ttl_are_suppressed(): + manager = make_manager() + entry = {"level_name": "INFO", "msg": "hello", "timestamp": 1000.0} + + await manager._handle_log("bot-1", entry) + await manager._handle_log("bot-1", dict(entry)) + + assert len(manager._bot_logs["bot-1"]) == 1 + assert len(manager._processed_messages) == 1 + + +async def test_error_logs_are_routed_to_the_error_deque(): + manager = make_manager() + + await manager._handle_log("bot-1", {"level_name": "ERROR", "msg": "boom", "timestamp": 1000.0}) + await manager._handle_log("bot-1", {"level_name": "INFO", "msg": "fine", "timestamp": 1000.0}) + await manager._handle_log("bot-1", "plain string log") + + assert [e["msg"] for e in manager._bot_error_logs["bot-1"]] == ["boom"] + assert [e["msg"] for e in manager._bot_logs["bot-1"]] == ["fine", "plain string log"] + + +async def test_cleanup_only_touches_the_expired_end_of_the_cache(): + manager = make_manager() + now = time.time() + + cache = CountingOrderedDict() + for i in range(3): + cache[f"expired-{i}"] = now - manager._message_ttl - 10 + for i in range(5000): + cache[f"live-{i}"] = now + manager._processed_messages = cache + + await manager._handle_log("bot-1", {"level_name": "INFO", "msg": "new", "timestamp": now}) + + # 3 pops for the expired entries + 1 peek that finds a live entry and stops. + assert cache.values_calls == 4 + assert not any(h.startswith("expired-") for h in cache) + assert len(cache) == 5001 + + +async def test_processed_messages_cache_is_bounded_regardless_of_log_rate(): + manager = make_manager() + manager._max_processed_messages = 50 + + for i in range(500): + await manager._handle_log("bot-1", {"level_name": "INFO", "msg": f"msg-{i}", "timestamp": 1000.0}) + + assert len(manager._processed_messages) == 50 + # The most recent messages are the ones retained. + assert "bot-1:msg-499:1000" in manager._processed_messages + assert "bot-1:msg-0:1000" not in manager._processed_messages diff --git a/test/test_order_sync_batches_queries.py b/test/test_order_sync_batches_queries.py new file mode 100644 index 00000000..0d22b9c6 --- /dev/null +++ b/test/test_order_sync_batches_queries.py @@ -0,0 +1,170 @@ +""" +Tests that `_sync_orders_to_database` reads the whole in-flight book with a single +batched SELECT instead of one round trip per order (PERF-049). + +Run with: pytest test/test_order_sync_batches_queries.py -v +""" +from contextlib import asynccontextmanager +from types import SimpleNamespace +from unittest.mock import MagicMock + +import pytest + +pytest.importorskip("hummingbot") + + +class FakeResult: + def __init__(self, rows): + self._rows = rows + + def scalars(self): + return self + + def all(self): + return list(self._rows) + + +class FakeSession: + """Records every statement executed and answers `client_order_id IN (...)` from memory.""" + + def __init__(self, rows): + self._rows_by_client_id = {row.client_order_id: row for row in rows} + self.statements = [] + self.flush_count = 0 + + async def execute(self, statement): + self.statements.append(statement) + requested = [] + for value in statement.compile().params.values(): + if isinstance(value, (list, tuple)): + requested.extend(value) + else: + requested.append(value) + return FakeResult( + [self._rows_by_client_id[cid] for cid in requested if cid in self._rows_by_client_id] + ) + + async def flush(self): + self.flush_count += 1 + + +def _service_with_session(session): + from services.unified_connector_service import UnifiedConnectorService + + @asynccontextmanager + async def get_session_context(): + yield session + + service = UnifiedConnectorService.__new__(UnifiedConnectorService) + service.db_manager = MagicMock() + service.db_manager.get_session_context = get_session_context + return service + + +def _connector_with_orders(states): + from hummingbot.core.data_type.in_flight_order import OrderState + + connector = MagicMock() + connector.in_flight_orders = { + client_order_id: SimpleNamespace(current_state=state or OrderState.OPEN) + for client_order_id, state in states.items() + } + return connector + + +def _db_row(client_order_id, status): + return SimpleNamespace(client_order_id=client_order_id, status=status) + + +class TestOrderSyncBatchesQueries: + @pytest.mark.asyncio + async def test_single_select_for_many_in_flight_orders(self): + """A book of several orders is read with one SELECT, not one per order.""" + from hummingbot.core.data_type.in_flight_order import OrderState + + client_order_ids = [f"OID-{i}" for i in range(8)] + session = FakeSession([_db_row(cid, "OPEN") for cid in client_order_ids]) + service = _service_with_session(session) + connector = _connector_with_orders({cid: OrderState.OPEN for cid in client_order_ids}) + + await service._sync_orders_to_database(connector, "master", "binance") + + assert len(session.statements) == 1 + # Nothing changed status, so nothing needed flushing either. + assert session.flush_count == 0 + + @pytest.mark.asyncio + async def test_status_changes_are_persisted_with_one_flush(self): + from hummingbot.core.data_type.in_flight_order import OrderState + + rows = [_db_row("OID-1", "SUBMITTED"), _db_row("OID-2", "OPEN")] + session = FakeSession(rows) + service = _service_with_session(session) + connector = _connector_with_orders({ + "OID-1": OrderState.OPEN, + "OID-2": OrderState.PARTIALLY_FILLED, + }) + + await service._sync_orders_to_database(connector, "master", "binance") + + assert rows[0].status == "OPEN" + assert rows[1].status == "PARTIALLY_FILLED" + assert len(session.statements) == 1 + assert session.flush_count == 1 + + @pytest.mark.asyncio + async def test_terminal_orders_are_popped_and_still_updated(self): + from hummingbot.core.data_type.in_flight_order import OrderState + + rows = [_db_row("OID-filled", "OPEN"), _db_row("OID-open", "OPEN")] + session = FakeSession(rows) + service = _service_with_session(session) + connector = _connector_with_orders({ + "OID-filled": OrderState.FILLED, + "OID-open": OrderState.OPEN, + }) + + await service._sync_orders_to_database(connector, "master", "binance") + + assert rows[0].status == "FILLED" + assert list(connector.in_flight_orders) == ["OID-open"] + + @pytest.mark.asyncio + async def test_orders_missing_from_the_database_are_skipped(self): + from hummingbot.core.data_type.in_flight_order import OrderState + + rows = [_db_row("OID-known", "OPEN")] + session = FakeSession(rows) + service = _service_with_session(session) + connector = _connector_with_orders({ + "OID-known": OrderState.PARTIALLY_FILLED, + "OID-unknown": OrderState.CANCELED, + }) + + await service._sync_orders_to_database(connector, "master", "binance") + + assert rows[0].status == "PARTIALLY_FILLED" + # The unknown order has no row to correct but is still terminal, so it is popped. + assert list(connector.in_flight_orders) == ["OID-known"] + + +class TestOrderRepositoryBatchLookup: + @pytest.mark.asyncio + async def test_no_query_for_an_empty_id_list(self): + from database.repositories.order_repository import OrderRepository + + session = FakeSession([]) + assert await OrderRepository(session).get_orders_by_client_ids([]) == [] + assert session.statements == [] + + @pytest.mark.asyncio + async def test_large_id_lists_are_chunked(self): + from database.repositories.order_repository import OrderRepository + + client_order_ids = [f"OID-{i}" for i in range(OrderRepository.CLIENT_ID_CHUNK_SIZE + 1)] + session = FakeSession([_db_row(cid, "OPEN") for cid in client_order_ids]) + + orders = await OrderRepository(session).get_orders_by_client_ids(client_order_ids) + + assert len(session.statements) == 2 + assert len(orders) == len(client_order_ids) diff --git a/test/test_performance_report_reports_db_failure.py b/test/test_performance_report_reports_db_failure.py new file mode 100644 index 00000000..f025807e --- /dev/null +++ b/test/test_performance_report_reports_db_failure.py @@ -0,0 +1,316 @@ +"""A database outage is reported as a failure, not as a performance report full of zeroes. + +`get_performance_report` builds a zeroed report and then fills it in from the database. +The whole database block used to sit inside a bare `except Exception` that logged and moved +on, so an unreachable database produced the untouched zeroed report: total_executors 0, +every PnL 0.0, win_rate 0.0. That is byte-identical to the report of an account that has +simply never run an executor, so no consumer could tell the two apart -- the route answered +200 with the zeroes and the `/ws/executors` performance channel pushed them to dashboards +as real numbers, with no error state and nothing marking them stale, for as long as the +outage lasted. + +The failure now propagates: the route turns it into a 500 and the push loop sends an +`error` frame on the channel. The zeroed report is left to mean exactly one thing -- an +empty dataset. + +Run with: pytest test/test_performance_report_reports_db_failure.py -v --asyncio-mode=auto +""" + +import asyncio +from contextlib import asynccontextmanager +from decimal import Decimal +from types import SimpleNamespace +from unittest.mock import AsyncMock, MagicMock + +import pytest +from sqlalchemy import create_engine +from sqlalchemy.orm import Session +from sqlalchemy.pool import StaticPool + +from database.models import ExecutorRecord +from services.executor_service import ExecutorService +from services.executor_ws_manager import ExecutorSubscription, ExecutorWebSocketManager + + +class DatabaseUnavailable(Exception): + """Stands in for whatever the driver raises when the database is unreachable.""" + + +class _AsyncSessionAdapter: + """The async surface the repository uses, over a real synchronous Session.""" + + def __init__(self, session: Session): + self._session = session + + async def execute(self, statement): + return self._session.execute(statement) + + async def commit(self): + self._session.commit() + + async def rollback(self): + self._session.rollback() + + async def close(self): + self._session.close() + + +@pytest.fixture +def db(): + """An in-memory executors table whose session factory can be made to fail.""" + engine = create_engine("sqlite://", connect_args={"check_same_thread": False}, + poolclass=StaticPool) + ExecutorRecord.__table__.create(engine) + + @asynccontextmanager + async def session_context(): + session = Session(engine) + adapter = _AsyncSessionAdapter(session) + try: + yield adapter + await adapter.commit() + except Exception: + await adapter.rollback() + raise + finally: + await adapter.close() + + def insert(pnls): + with Session(engine) as session: + for i, pnl in enumerate(pnls): + session.add(ExecutorRecord( + executor_id=f"e-{i}", + executor_type="position_executor", + account_name="master", + connector_name="binance_perpetual", + trading_pair="BTC-USDT", + controller_id="main", + status="TERMINATED", + close_type="TAKE_PROFIT", + net_pnl_quote=Decimal(str(pnl)), + net_pnl_pct=Decimal("0"), + cum_fees_quote=Decimal("1"), + filled_amount_quote=Decimal("100"), + )) + session.commit() + + try: + yield SimpleNamespace(session_context=session_context, insert=insert) + finally: + engine.dispose() + + +def _service(session_context): + service = ExecutorService.__new__(ExecutorService) + service.db_manager = MagicMock(get_session_context=session_context) + service._executor_metadata = {} + service._active_executors = {} + service._positions_held = {} + return service + + +@asynccontextmanager +async def _unreachable_database(): + """A session context that fails the way a down database does.""" + raise DatabaseUnavailable("could not connect to server") + yield # pragma: no cover - unreachable, keeps this an async generator + + +# -------------------------------------------------------------------------------------- +# The service +# -------------------------------------------------------------------------------------- + +class TestTheService: + + async def test_a_database_outage_raises_instead_of_reporting_zeroes(self): + """The bug: the outage was swallowed and its zeroed report returned as data.""" + service = _service(lambda: _unreachable_database()) + + with pytest.raises(DatabaseUnavailable): + await service.get_performance_report() + + async def test_a_failure_midway_through_the_query_also_raises(self, db): + """Not only the connect: a query that dies partway through must surface too.""" + service = _service(db.session_context) + service.db_manager.get_session_context = lambda: _unreachable_database() + + with pytest.raises(DatabaseUnavailable): + await service.get_performance_report(controller_id="main") + + async def test_an_empty_dataset_still_reports_zeroes(self, db): + """The other half: zeroes must stay the honest answer for an empty table.""" + report = await _service(db.session_context).get_performance_report() + + assert report["total_executors"] == 0 + assert report["by_status"] == {} + assert report["pnl_total_quote"] == 0.0 + assert report["global_pnl_quote"] == 0.0 + assert report["win_rate"] == 0.0 + assert report["sharpe_ratio"] is None + assert report["by_type"] == [] + + async def test_a_populated_dataset_is_unaffected(self, db): + """Removing the except must not change the report the database can answer.""" + db.insert([10.0, -4.0, 6.0]) + + report = await _service(db.session_context).get_performance_report() + + assert report["total_executors"] == 3 + assert report["pnl_total_quote"] == pytest.approx(12.0) + assert report["volume_total_quote"] == pytest.approx(300.0) + assert report["fees_total_quote"] == pytest.approx(3.0) + assert report["win_rate"] == pytest.approx(2 / 3) + + +# -------------------------------------------------------------------------------------- +# The route +# -------------------------------------------------------------------------------------- + +class TestTheRoute: + + def _client(self, report): + from fastapi import FastAPI + from fastapi.testclient import TestClient + + import routers.executors as executors_router + from deps import get_executor_service, get_market_data_service + + executor_service = MagicMock() + executor_service.get_performance_report = report + + app = FastAPI() + app.include_router(executors_router.router) + app.dependency_overrides[get_executor_service] = lambda: executor_service + app.dependency_overrides[get_market_data_service] = lambda: MagicMock() + return TestClient(app, raise_server_exceptions=False) + + def test_an_outage_answers_500_rather_than_200_with_zeroes(self): + client = self._client(AsyncMock(side_effect=DatabaseUnavailable("no server"))) + + response = client.get("/executors/performance") + + assert response.status_code == 500 + + def test_an_empty_dataset_answers_200_with_zeroes(self): + empty = { + "controller_id": None, "total_executors": 0, "by_status": {}, + "pnl_total_quote": 0.0, "unrealized_pnl_quote": 0.0, "global_pnl_quote": 0.0, + "pnl_pct_avg": 0.0, "fees_total_quote": 0.0, "volume_total_quote": 0.0, + "win_rate": 0.0, "sharpe_ratio": None, "by_type": [], "active_positions": 0, + } + client = self._client(AsyncMock(return_value=empty)) + + response = client.get("/executors/performance") + + assert response.status_code == 200 + assert response.json()["total_executors"] == 0 + + +# -------------------------------------------------------------------------------------- +# The WebSocket performance channel +# -------------------------------------------------------------------------------------- + +class RecordingWebSocket: + """Captures pushed frames and lets a test await the Nth one.""" + + def __init__(self): + self.sent = [] + self._target = None + self._reached = asyncio.Event() + + async def send_json(self, message): + self.sent.append(message) + if self._target is not None and len(self.sent) >= self._target: + self._reached.set() + + async def wait_for_frames(self, count, timeout=2.0): + self._target = count + if len(self.sent) >= count: + return + self._reached.clear() + await asyncio.wait_for(self._reached.wait(), timeout) + + +def _manager(get_performance_report): + executor_service = MagicMock() + executor_service.get_performance_report = get_performance_report + return ExecutorWebSocketManager( + executor_service=executor_service, + market_data_service=MagicMock(), + bots_orchestrator=MagicMock(), + ) + + +async def _run(manager, sub, websocket, frames): + task = asyncio.create_task( + manager._get_push_fn("performance")("conn-1", websocket, sub) + ) + try: + await websocket.wait_for_frames(frames) + finally: + task.cancel() + await asyncio.gather(task, return_exceptions=True) + + +def _sub(): + return ExecutorSubscription( + sub_id="sub-1", sub_type="performance", update_interval=0.01, controller_id="main" + ) + + +class TestTheWebSocketChannel: + + async def test_an_outage_pushes_an_error_frame_naming_the_channel(self): + manager = _manager(AsyncMock(side_effect=DatabaseUnavailable("no server"))) + websocket = RecordingWebSocket() + + await _run(manager, _sub(), websocket, frames=1) + + frame = websocket.sent[0] + assert frame["type"] == "error" + assert frame["channel"] == "performance" + assert frame["subscription_id"] == "sub-1" + assert "no server" in frame["message"] + + async def test_a_sustained_outage_sends_one_error_frame_not_one_per_interval(self): + """The client is told once; the loop keeps retrying quietly behind it.""" + manager = _manager(AsyncMock(side_effect=DatabaseUnavailable("no server"))) + websocket = RecordingWebSocket() + sub = _sub() + + task = asyncio.create_task( + manager._get_push_fn("performance")("conn-1", websocket, sub) + ) + try: + await websocket.wait_for_frames(1) + await asyncio.sleep(0.1) # ~10 more failing polls + finally: + task.cancel() + await asyncio.gather(task, return_exceptions=True) + + assert len(websocket.sent) == 1 + + async def test_the_channel_recovers_with_a_data_frame_after_the_error(self): + """And the recovery frame is sent even though the payload never changed.""" + report = {"total_executors": 0, "pnl_total_quote": 0.0} + manager = _manager(AsyncMock(side_effect=[ + report, DatabaseUnavailable("no server"), report, report, + ])) + websocket = RecordingWebSocket() + + await _run(manager, _sub(), websocket, frames=3) + + types = [f["type"] for f in websocket.sent[:3]] + assert types == ["performance", "error", "performance"] + assert websocket.sent[2]["data"] == report + + async def test_an_empty_report_is_still_pushed_as_data(self): + """A genuinely empty dataset is not an error and must reach the client.""" + empty = {"total_executors": 0, "pnl_total_quote": 0.0, "win_rate": 0.0} + manager = _manager(AsyncMock(return_value=empty)) + websocket = RecordingWebSocket() + + await _run(manager, _sub(), websocket, frames=1) + + assert websocket.sent[0]["type"] == "performance" + assert websocket.sent[0]["data"] == empty diff --git a/test/test_performance_report_sharpe_is_aggregated.py b/test/test_performance_report_sharpe_is_aggregated.py new file mode 100644 index 00000000..ab4d438a --- /dev/null +++ b/test/test_performance_report_sharpe_is_aggregated.py @@ -0,0 +1,248 @@ +"""The Sharpe ratio is computed by the database, not by shipping every PnL row to Python. + +`get_performance_report` used to run a second, unfiltered `SELECT net_pnl_quote` over every +completed executor and hand the list to `ExecutorService`, which took mean and variance in +Python. The list had exactly one consumer -- the Sharpe ratio -- and no LIMIT, so the report +scanned and transferred the whole executors table. That report is polled by the +`/ws/executors` performance push loop every `update_interval` seconds *per subscriber*, so +the cost of a client watching a chart grew without bound as the table grew. + +It is now one more aggregate in the query that already computes sum/avg/count/win-rate over +the same filter: the count, the sum and the sum of squares give the sample standard +deviation directly. What is pinned here is that the row count is back to O(1), and that the +number the API reports is still the number the per-row computation produced. +""" + +import inspect +import math +import re +from contextlib import asynccontextmanager +from decimal import Decimal +from types import SimpleNamespace +from unittest.mock import MagicMock + +import pytest +from sqlalchemy import create_engine +from sqlalchemy.orm import Session +from sqlalchemy.pool import StaticPool + +from database.models import ExecutorRecord +from database.repositories.executor_repository import ExecutorRepository +from services.executor_service import ExecutorService + + +class _AsyncSessionAdapter: + """The async surface the repository uses, over a real synchronous Session. + + aiosqlite is not installed here and the report only ever awaits `execute`, so this + runs the real SQL against the real schema. Mocking the session would prove nothing: + the whole point is which statements the database is asked to run. + """ + + def __init__(self, session: Session): + self._session = session + self.statements = [] + + async def execute(self, statement): + self.statements.append(statement) + return self._session.execute(statement) + + async def commit(self): + self._session.commit() + + async def rollback(self): + self._session.rollback() + + async def close(self): + self._session.close() + + +@pytest.fixture +def db(): + """An in-memory executors table, plus a session factory that records its SQL.""" + engine = create_engine("sqlite://", connect_args={"check_same_thread": False}, + poolclass=StaticPool) + ExecutorRecord.__table__.create(engine) + + sessions = [] + + @asynccontextmanager + async def session_context(): + session = Session(engine) + adapter = _AsyncSessionAdapter(session) + sessions.append(adapter) + try: + yield adapter + await adapter.commit() + except Exception: + await adapter.rollback() + raise + finally: + await adapter.close() + + def insert(pnls, controller_id="main"): + with Session(engine) as session: + for i, pnl in enumerate(pnls): + session.add(ExecutorRecord( + executor_id=f"e-{controller_id}-{i}", + executor_type="position_executor", + account_name="master", + connector_name="binance_perpetual", + trading_pair="BTC-USDT", + controller_id=controller_id, + status="TERMINATED", + close_type="TAKE_PROFIT", + net_pnl_quote=None if pnl is None else Decimal(str(pnl)), + net_pnl_pct=Decimal("0"), + cum_fees_quote=Decimal("0"), + filled_amount_quote=Decimal("100"), + )) + session.commit() + + try: + yield SimpleNamespace(session_context=session_context, insert=insert, + sessions=sessions, engine=engine) + finally: + engine.dispose() + + +def _sharpe_the_old_way(pnls): + """The exact Python the service used to run over the fetched per-executor rows.""" + values = [float(p or 0) for p in pnls] + if len(values) < 2: + return None + mean = sum(values) / len(values) + variance = sum((v - mean) ** 2 for v in values) / (len(values) - 1) + std = math.sqrt(variance) + return round(mean / std, 4) if std > 0 else None + + +def _service(db): + service = ExecutorService.__new__(ExecutorService) + service.db_manager = MagicMock(get_session_context=db.session_context) + service._executor_metadata = {} + service._active_executors = {} + service._positions_held = {} + return service + + +async def _report(db, controller_id=None): + return await _service(db).get_performance_report(controller_id=controller_id) + + +# -------------------------------------------------------------------------------------- +# The row count no longer follows the table size +# -------------------------------------------------------------------------------------- + +def test_the_report_never_selects_a_bare_pnl_column(): + """A `select(ExecutorRecord.net_pnl_quote)` here is a full scan with no LIMIT.""" + source = inspect.getsource(ExecutorRepository.get_performance_report) + selected = re.findall(r"select\(\s*ExecutorRecord\.(\w+)\b", source) + + assert "net_pnl_quote" not in selected, ( + "the report is fetching one PnL row per completed executor again; " + "the Sharpe inputs belong in the aggregate query" + ) + + +@pytest.mark.asyncio +async def test_the_report_costs_the_same_number_of_rows_at_any_table_size(db): + """Ten executors and a thousand must return the same number of rows to Python.""" + db.insert([1.0, -2.0, 3.0, -0.5, 4.25, -1.75, 0.5, 2.0, -3.0, 1.25]) + async with db.session_context() as session: + small = await ExecutorRepository(session).get_performance_report() + + db.insert([i * 0.01 - 5 for i in range(1000)], controller_id="bulk") + async with db.session_context() as session: + big = await ExecutorRepository(session).get_performance_report() + + # The payload is aggregates only: no per-executor sequence of any kind. + for report in (small, big): + assert not any(isinstance(v, (list, tuple)) and v and isinstance(v[0], float) + for v in report.values()), f"a per-row list leaked back in: {report}" + + def db_rows(report): + """Rows crossing the wire: the aggregate row, plus one per group.""" + return 1 + len(report["status_counts"]) + len(report["by_type"]) + + assert db_rows(big) == db_rows(small), ( + "the report got more expensive purely because the table got bigger" + ) + + +# -------------------------------------------------------------------------------------- +# The number itself did not change +# -------------------------------------------------------------------------------------- + +@pytest.mark.asyncio +@pytest.mark.parametrize("pnls", [ + [1.0, -2.0, 3.0, -0.5, 4.25, -1.75, 0.5, 2.0, -3.0, 1.25], # mixed + [10.0, 12.0], # the two-row minimum + [-4.0, -1.0, -9.0, -2.5], # a losing controller + [1.0, None, 3.0, -2.0], # NULL PnL counts as zero + [0.001, 0.002, 0.0015, 0.0011], # tiny, tightly clustered + [1_000_000.5, 1_000_001.5, 1_000_000.0], # large mean, small spread +]) +async def test_the_sharpe_ratio_matches_the_per_row_computation(db, pnls): + db.insert(pnls) + report = await _report(db) + + assert report["sharpe_ratio"] == _sharpe_the_old_way(pnls) + assert report["sharpe_ratio"] is not None + + +@pytest.mark.asyncio +async def test_a_single_executor_has_no_sharpe_ratio(db): + """stddev of one sample is undefined -- the old guard was `len(pnl_values) >= 2`.""" + db.insert([7.5]) + report = await _report(db) + + assert report["sharpe_ratio"] is None + assert report["total_executors"] == 1 + + +@pytest.mark.asyncio +async def test_no_executors_at_all_has_no_sharpe_ratio(db): + report = await _report(db) + + assert report["sharpe_ratio"] is None + assert report["pnl_total_quote"] == 0.0 + + +@pytest.mark.asyncio +async def test_identical_pnls_have_no_sharpe_ratio_rather_than_dividing_by_zero(db): + """Zero variance: float error around the subtraction must not become a real number.""" + db.insert([2.5, 2.5, 2.5, 2.5]) + report = await _report(db) + + assert report["sharpe_ratio"] is None + + +@pytest.mark.asyncio +async def test_a_position_hold_is_left_out_of_the_sharpe_ratio_too(db): + """The dispersion aggregate rides the same filter as the PnL total it belongs to.""" + db.insert([1.0, -2.0, 3.0, -0.5]) + with Session(db.engine) as session: + session.add(ExecutorRecord( + executor_id="e-hold", executor_type="position_executor", account_name="master", + connector_name="binance_perpetual", trading_pair="BTC-USDT", controller_id="main", + status="TERMINATED", close_type="POSITION_HOLD", + net_pnl_quote=Decimal("500"), net_pnl_pct=Decimal("0"), + cum_fees_quote=Decimal("0"), filled_amount_quote=Decimal("100"), + )) + session.commit() + + report = await _report(db) + + assert report["sharpe_ratio"] == _sharpe_the_old_way([1.0, -2.0, 3.0, -0.5]) + + +@pytest.mark.asyncio +async def test_the_controller_filter_still_narrows_the_sharpe_ratio(db): + mine = [1.0, -2.0, 3.0, -0.5] + db.insert(mine, controller_id="mine") + db.insert([100.0, -100.0, 250.0], controller_id="theirs") + + report = await _report(db, controller_id="mine") + + assert report["sharpe_ratio"] == _sharpe_the_old_way(mine) diff --git a/test/test_poller_writes_through_the_services.py b/test/test_poller_writes_through_the_services.py new file mode 100644 index 00000000..9e24ede3 --- /dev/null +++ b/test/test_poller_writes_through_the_services.py @@ -0,0 +1,217 @@ +"""The poller decides what the chain says; the services decide what gets written. + +The poller used to build `GatewayCLMMRepository` / `GatewaySwapRepository` itself in six +places and hold one session open across a whole cycle of Gateway calls (ARCH-103). It now +reads its work as plain dicts and writes through the same services the `/gateway/clmm/*` +and `/gateway/swap*` routes persist through. + +Every call site changed, so these pin the decisions that must have survived the move: a +transaction is only aged out after a poll that actually reached the chain, a NOT_FOUND is +given its grace window before it counts as dropped, and a position is only closed after +the consecutive-miss gate says so. +""" +from datetime import datetime, timedelta, timezone +from types import SimpleNamespace +from unittest.mock import AsyncMock + +import pytest + +from services.gateway_transaction_poller import GatewayTransactionPoller + + +def _poller(poll_result=None, position_info=None): + """A poller with both services stubbed, so writes are observable as calls.""" + poller = object.__new__(GatewayTransactionPoller) + poller.max_retry_age = 3600 + poller._position_missing_strikes = {} + poller.gateway_client = SimpleNamespace( + ping=AsyncMock(return_value=True), + poll_transaction=AsyncMock(return_value=poll_result), + clmm_position_info=AsyncMock(return_value=position_info), + ) + poller.swap_service = SimpleNamespace(update_swap_status=AsyncMock()) + poller.clmm_service = SimpleNamespace( + record_event_confirmed=AsyncMock(), + update_event_status=AsyncMock(), + mark_position_closed=AsyncMock(), + record_position_state=AsyncMock(), + ) + return poller + + +def _swap(age_seconds=10): + return { + "transaction_hash": "TX", + "network": "solana-mainnet-beta", + "timestamp": datetime.now(timezone.utc) - timedelta(seconds=age_seconds), + } + + +def _event(age_seconds=10, network="solana-mainnet-beta"): + return { + "transaction_hash": "TX", + "network": network, + "position_address": "POS", + "timestamp": datetime.now(timezone.utc) - timedelta(seconds=age_seconds), + } + + +def _position(): + return { + "id": 7, + "position_address": "POS", + "wallet_address": "WALLET", + "connector": "meteora", + "network": "solana-mainnet-beta", + } + + +# --------------------------------------------------------------------------- +# Swaps +# --------------------------------------------------------------------------- + +@pytest.mark.asyncio +async def test_a_confirmed_swap_is_recorded_with_its_gas(): + poller = _poller({"txStatus": 1, "fee": 0.000005}) + await poller._poll_swap_transaction(_swap()) + + written = poller.swap_service.update_swap_status.await_args.kwargs + assert written["transaction_hash"] == "TX" + assert written["status"] == "CONFIRMED" + assert float(written["gas_fee"]) == 0.000005 + assert written["gas_token"] == "SOL" + + +@pytest.mark.asyncio +async def test_a_transient_gateway_error_writes_nothing(): + # No information came back, so no state may change — the swap is polled again. + poller = _poller({"error": "Gateway 500", "status": 500}) + await poller._poll_swap_transaction(_swap(age_seconds=99999)) + poller.swap_service.update_swap_status.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_not_found_is_only_dropped_once_the_blockhash_can_no_longer_be_valid(): + poller = _poller({"txStatus": -2}) + await poller._poll_swap_transaction(_swap(age_seconds=30)) + poller.swap_service.update_swap_status.assert_not_awaited() + + await poller._poll_swap_transaction(_swap(age_seconds=600)) + assert poller.swap_service.update_swap_status.await_args.kwargs["status"] == "FAILED" + + +@pytest.mark.asyncio +async def test_a_swap_still_pending_past_the_retry_age_is_timed_out(): + poller = _poller({"txStatus": 0}) + await poller._poll_swap_transaction(_swap(age_seconds=60)) + poller.swap_service.update_swap_status.assert_not_awaited() + + await poller._poll_swap_transaction(_swap(age_seconds=7200)) + written = poller.swap_service.update_swap_status.await_args.kwargs + assert written["status"] == "FAILED" + assert written["error_message"] == "Transaction confirmation timeout" + + +# --------------------------------------------------------------------------- +# CLMM events +# --------------------------------------------------------------------------- + +@pytest.mark.asyncio +async def test_a_confirmed_event_is_booked_against_its_position(): + # One service call: the status and the position bookkeeping it owes share a session. + poller = _poller({"txStatus": 1, "fee": 0.000011772}) + await poller._poll_clmm_event_transaction(_event()) + + written = poller.clmm_service.record_event_confirmed.await_args.kwargs + assert written["transaction_hash"] == "TX" + assert float(written["gas_fee"]) == 0.000011772 + poller.clmm_service.update_event_status.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_a_failed_event_records_the_reason_and_books_nothing(): + poller = _poller({"txStatus": -1, "error": "SLIPPAGE_EXCEEDED (0x1771)", "fee": 0.000005}) + await poller._poll_clmm_event_transaction(_event()) + + written = poller.clmm_service.update_event_status.await_args.kwargs + assert written["status"] == "FAILED" + assert "SLIPPAGE_EXCEEDED" in written["error_message"] + poller.clmm_service.record_event_confirmed.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_an_event_whose_position_row_is_missing_is_not_polled(): + # Without the position there is no chain to ask about the event. + poller = _poller({"txStatus": 1}) + await poller._poll_clmm_event_transaction(_event(network=None)) + + poller.gateway_client.poll_transaction.assert_not_awaited() + poller.clmm_service.record_event_confirmed.assert_not_awaited() + + +# --------------------------------------------------------------------------- +# Position state +# --------------------------------------------------------------------------- + +@pytest.mark.asyncio +async def test_a_live_position_writes_back_one_reading(): + poller = _poller(position_info={ + "address": "POS", + "price": 200.0, + "lowerPrice": 150.0, + "upperPrice": 250.0, + "baseTokenAmount": 0.0099, + "quoteTokenAmount": 1.98, + "baseFeeAmount": 0.00031, + "quoteFeeAmount": 0.062, + }) + await poller._refresh_position_state(_position()) + + written = poller.clmm_service.record_position_state.await_args.kwargs + assert written["position_address"] == "POS" + assert written["in_range"] == "IN_RANGE" + assert float(written["current_price"]) == 200.0 + assert float(written["base_fee_pending"]) == 0.00031 + poller.clmm_service.mark_position_closed.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_zero_liquidity_closes_the_position(): + poller = _poller(position_info={ + "address": "POS", "price": 200.0, "lowerPrice": 150.0, "upperPrice": 250.0, + "baseTokenAmount": 0, "quoteTokenAmount": 0, + }) + await poller._refresh_position_state(_position()) + + poller.clmm_service.mark_position_closed.assert_awaited_once_with("POS") + poller.clmm_service.record_position_state.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_a_position_closes_only_after_three_consecutive_misses(): + # Gateway 500s on transient RPC trouble as well as on a position that is gone, + # so one miss must never close a live position. + poller = _poller(position_info={"error": "not found", "status": 500}) + + for _ in range(GatewayTransactionPoller.MISSING_STRIKES_TO_CLOSE - 1): + await poller._refresh_position_state(_position()) + poller.clmm_service.mark_position_closed.assert_not_awaited() + + await poller._refresh_position_state(_position()) + poller.clmm_service.mark_position_closed.assert_awaited_once_with("POS") + # The strike count is cleared with the position it belonged to. + assert "POS" not in poller._position_missing_strikes + + +@pytest.mark.asyncio +async def test_a_successful_read_clears_earlier_misses(): + poller = _poller(position_info={"error": "not found", "status": 500}) + await poller._refresh_position_state(_position()) + assert poller._position_missing_strikes["POS"] == 1 + + poller.gateway_client.clmm_position_info = AsyncMock(return_value={ + "address": "POS", "price": 200.0, "lowerPrice": 150.0, "upperPrice": 250.0, + "baseTokenAmount": 0.0099, "quoteTokenAmount": 1.98, + }) + await poller._refresh_position_state(_position()) + assert "POS" not in poller._position_missing_strikes diff --git a/test/test_reconcile_batches_queries.py b/test/test_reconcile_batches_queries.py new file mode 100644 index 00000000..534d19f5 --- /dev/null +++ b/test/test_reconcile_batches_queries.py @@ -0,0 +1,247 @@ +""" +Tests that `reconcile_active_orders` reads the tracked book with one batched SELECT and +flushes once, instead of a SELECT plus a flush per order (PERF-108). + +Run with: pytest test/test_reconcile_batches_queries.py -v +""" +from contextlib import asynccontextmanager +from types import SimpleNamespace +from unittest.mock import MagicMock + +import pytest + +pytest.importorskip("hummingbot") + + +class FakeResult: + def __init__(self, rows): + self._rows = rows + + def scalars(self): + return self + + def all(self): + return list(self._rows) + + +class FakeSession: + """Records every statement executed and answers `client_order_id IN (...)` from memory.""" + + def __init__(self, rows, flush_error=None): + self._rows_by_client_id = {row.client_order_id: row for row in rows} + self._flush_error = flush_error + self.statements = [] + self.flush_count = 0 + + async def execute(self, statement): + self.statements.append(statement) + requested = [] + for value in statement.compile().params.values(): + if isinstance(value, (list, tuple)): + requested.extend(value) + else: + requested.append(value) + return FakeResult( + [self._rows_by_client_id[cid] for cid in requested if cid in self._rows_by_client_id] + ) + + async def flush(self): + self.flush_count += 1 + if self._flush_error is not None: + raise self._flush_error + + +def _service_with_session(session, connector, account_name="master", connector_name="binance"): + from services.unified_connector_service import UnifiedConnectorService + + @asynccontextmanager + async def get_session_context(): + yield session + + service = UnifiedConnectorService.__new__(UnifiedConnectorService) + service.db_manager = MagicMock() + service.db_manager.get_session_context = get_session_context + service._trading_connectors = {account_name: {connector_name: connector}} + return service + + +class OrderNotFound(Exception): + """Stands in for the connector-specific "unknown order" error.""" + + +def _connector_with_orders(exchange_states): + """Connector whose `_request_order_status` answers with `exchange_states[client_order_id]`. + + A value that is an exception instance is raised instead of returned. + """ + connector = MagicMock() + connector.in_flight_orders = { + client_order_id: SimpleNamespace(client_order_id=client_order_id) + for client_order_id in exchange_states + } + + async def _request_order_status(order): + outcome = exchange_states[order.client_order_id] + if isinstance(outcome, Exception): + raise outcome + return SimpleNamespace(new_state=outcome) + + connector._request_order_status = _request_order_status + connector._is_order_not_found_during_status_update_error = ( + lambda exc: isinstance(exc, OrderNotFound) + ) + return connector + + +def _db_row(client_order_id, status, error_message=None): + return SimpleNamespace( + client_order_id=client_order_id, status=status, error_message=error_message + ) + + +class TestReconcileBatchesQueries: + @pytest.mark.asyncio + async def test_query_count_does_not_grow_with_the_order_count(self): + """The number of SELECTs depends on the chunk size, never on the order count.""" + from hummingbot.core.data_type.in_flight_order import OrderState + + from database.repositories.order_repository import OrderRepository + + statement_counts = [] + for order_count in (1, 8, 64): + client_order_ids = [f"OID-{i}" for i in range(order_count)] + session = FakeSession([_db_row(cid, "OPEN") for cid in client_order_ids]) + connector = _connector_with_orders({cid: OrderState.OPEN for cid in client_order_ids}) + service = _service_with_session(session, connector) + + summary = await service.reconcile_active_orders() + + assert summary["still_open"] == order_count + statement_counts.append(len(session.statements)) + + assert statement_counts == [1, 1, 1] + # And the batching is what keeps it flat: a book larger than one chunk still costs + # ceil(N / CLIENT_ID_CHUNK_SIZE) queries, not N. + assert OrderRepository.CLIENT_ID_CHUNK_SIZE > 1 + + @pytest.mark.asyncio + async def test_large_book_costs_one_query_per_chunk(self): + from hummingbot.core.data_type.in_flight_order import OrderState + + from database.repositories.order_repository import OrderRepository + + order_count = OrderRepository.CLIENT_ID_CHUNK_SIZE + 1 + client_order_ids = [f"OID-{i}" for i in range(order_count)] + session = FakeSession([_db_row(cid, "OPEN") for cid in client_order_ids]) + connector = _connector_with_orders({cid: OrderState.OPEN for cid in client_order_ids}) + service = _service_with_session(session, connector) + + await service.reconcile_active_orders() + + assert len(session.statements) == 2 + + @pytest.mark.asyncio + async def test_no_flush_when_no_status_changed(self): + from hummingbot.core.data_type.in_flight_order import OrderState + + client_order_ids = [f"OID-{i}" for i in range(5)] + session = FakeSession([_db_row(cid, "OPEN") for cid in client_order_ids]) + connector = _connector_with_orders({cid: OrderState.OPEN for cid in client_order_ids}) + service = _service_with_session(session, connector) + + await service.reconcile_active_orders() + + assert session.flush_count == 0 + + @pytest.mark.asyncio + async def test_status_changes_are_persisted_with_a_single_flush(self): + from hummingbot.core.data_type.in_flight_order import OrderState + + rows = [ + _db_row("OID-filled", "OPEN"), + _db_row("OID-open", "SUBMITTED"), + _db_row("OID-gone", "OPEN"), + ] + session = FakeSession(rows) + connector = _connector_with_orders({ + "OID-filled": OrderState.FILLED, + "OID-open": OrderState.OPEN, + "OID-gone": OrderNotFound("unknown order"), + }) + service = _service_with_session(session, connector) + + summary = await service.reconcile_active_orders() + + assert session.flush_count == 1 + assert rows[0].status == "FILLED" + assert rows[1].status == "OPEN" + assert rows[2].status == "CANCELLED" + assert rows[2].error_message == "Reconciled on startup: order not found on exchange" + assert summary["reconciled_terminal"] == 2 + assert summary["still_open"] == 1 + # Terminal orders stop being tracked; the open one stays cancelable. + assert list(connector.in_flight_orders) == ["OID-open"] + + @pytest.mark.asyncio + async def test_orders_missing_from_the_database_are_skipped(self): + from hummingbot.core.data_type.in_flight_order import OrderState + + rows = [_db_row("OID-known", "OPEN")] + session = FakeSession(rows) + connector = _connector_with_orders({ + "OID-known": OrderState.PARTIALLY_FILLED, + "OID-unknown": OrderState.FILLED, + }) + service = _service_with_session(session, connector) + + summary = await service.reconcile_active_orders() + + assert rows[0].status == "PARTIALLY_FILLED" + # The order with no row is still reconciled against the exchange and untracked. + assert summary["reconciled_terminal"] == 1 + assert summary["still_open"] == 1 + assert list(connector.in_flight_orders) == ["OID-known"] + + @pytest.mark.asyncio + async def test_unverifiable_orders_are_left_untouched(self): + from hummingbot.core.data_type.in_flight_order import OrderState + + rows = [_db_row("OID-flaky", "OPEN"), _db_row("OID-ok", "SUBMITTED")] + session = FakeSession(rows) + connector = _connector_with_orders({ + "OID-flaky": TimeoutError("exchange unreachable"), + "OID-ok": OrderState.OPEN, + }) + service = _service_with_session(session, connector) + + summary = await service.reconcile_active_orders() + + assert summary["unverified"] == 1 + assert rows[0].status == "OPEN" + assert rows[1].status == "OPEN" + assert session.flush_count == 1 + # Nothing terminal was confirmed, so the whole book stays tracked. + assert list(connector.in_flight_orders) == ["OID-flaky", "OID-ok"] + + @pytest.mark.asyncio + async def test_a_failed_flush_leaves_tracking_untouched(self): + from hummingbot.core.data_type.in_flight_order import OrderState + + rows = [_db_row("OID-filled", "OPEN"), _db_row("OID-open", "SUBMITTED")] + session = FakeSession(rows, flush_error=RuntimeError("db down")) + connector = _connector_with_orders({ + "OID-filled": OrderState.FILLED, + "OID-open": OrderState.OPEN, + }) + service = _service_with_session(session, connector) + + summary = await service.reconcile_active_orders() + + assert summary == { + "reconciled_terminal": 0, + "still_open": 0, + "unverified": 2, + "skipped_connectors": 0, + } + # Nothing was untracked, so the next startup reconciles these orders again. + assert list(connector.in_flight_orders) == ["OID-filled", "OID-open"] diff --git a/test/test_remove_container_is_api_managed.py b/test/test_remove_container_is_api_managed.py new file mode 100644 index 00000000..a3fa6784 --- /dev/null +++ b/test/test_remove_container_is_api_managed.py @@ -0,0 +1,98 @@ +""" +Tests for the ownership guard on ``POST /docker/remove-container/{container_name}``. + +The endpoint used to refuse any name that did not start with ``hummingbot-``. Nothing names bot +containers that way — ``DockerService.create_hummingbot_instance`` passes the instance name to +Docker verbatim — so the guard rejected exactly the containers this API creates while happily +accepting the infrastructure containers that *do* carry the prefix (``hummingbot-postgres``, +``hummingbot-broker``). + +The guard is now the real invariant: a container is removable here only if it owns the +``bots/instances/`` directory that this endpoint archives. + +Run with: pytest test/test_remove_container_is_api_managed.py -v +""" +import os + +import pytest +from fastapi import HTTPException + +from routers.docker import remove_container + + +class _StubDocker: + def __init__(self): + self.removed = [] + + def remove_container(self, container_name): + self.removed.append(container_name) + return {"success": True, "message": f"Container {container_name} removed successfully."} + + +class _StubArchiver: + def __init__(self): + self.archived = [] + + def archive_locally(self, instance_name, instance_dir): + self.archived.append((instance_name, instance_dir)) + + +@pytest.fixture +def instances_root(tmp_path, monkeypatch): + """Run inside a scratch cwd so 'bots/instances' is a real, empty directory we control.""" + root = tmp_path / "bots" / "instances" + root.mkdir(parents=True) + monkeypatch.chdir(tmp_path) + return root + + +async def _remove(name, docker, archiver): + return await remove_container( + container_name=name, + archive_locally=True, + s3_bucket=None, + docker_service=docker, + bot_archiver=archiver, + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("bot_name", ["pmm-1", "hummingbot-pmm-1", "my_bot.v2"]) +async def test_api_managed_container_is_removed_whatever_its_name(instances_root, bot_name): + """A container this API created is removable regardless of whether its name carries a prefix.""" + (instances_root / bot_name).mkdir() + docker, archiver = _StubDocker(), _StubArchiver() + + response = await _remove(bot_name, docker, archiver) + + assert response["success"] is True + assert docker.removed == [bot_name] + assert archiver.archived == [(bot_name, os.path.join("bots", "instances", bot_name))] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("container_name", ["nginx", "hummingbot-postgres", "hummingbot-broker"]) +async def test_unrelated_host_container_is_refused(instances_root, container_name): + """No instance directory means it is not ours - including the prefixed infra containers.""" + docker, archiver = _StubDocker(), _StubArchiver() + + with pytest.raises(HTTPException) as exc: + await _remove(container_name, docker, archiver) + + assert exc.value.status_code == 400 + assert "managed by this API" in exc.value.detail + assert docker.removed == [] + assert archiver.archived == [] + + +@pytest.mark.asyncio +async def test_name_escaping_the_instances_directory_is_refused(instances_root): + """Containment still holds: a traversing name never reaches remove/archive.""" + docker, archiver = _StubDocker(), _StubArchiver() + + with pytest.raises(HTTPException) as exc: + await _remove(os.path.join("..", "credentials"), docker, archiver) + + assert exc.value.status_code == 400 + assert docker.removed == [] + assert archiver.archived == [] diff --git a/test/test_stop_and_archive_bot_name.py b/test/test_stop_and_archive_bot_name.py new file mode 100644 index 00000000..c3dc5ef4 --- /dev/null +++ b/test/test_stop_and_archive_bot_name.py @@ -0,0 +1,213 @@ +""" +Tests for the stop-and-archive bot-name plumbing. + +``POST /stop-and-archive-bot/{bot_name}`` used to thread three aliases of the +same string (``actual_bot_name`` / ``container_name`` / +``bot_name_for_orchestrator``) into the background task. They are now a single +``bot_name``; these tests pin that the name actually used for every +destructive step (stop container, archive directory, remove container) is the +path parameter, unmodified. + +Run with: pytest test/test_stop_and_archive_bot_name.py -v +""" +import os +from contextlib import asynccontextmanager +from types import SimpleNamespace + +import pytest +from fastapi import BackgroundTasks + +import services.bots_orchestrator as orchestrator_module +from routers.bot_orchestration import stop_and_archive_bot as stop_and_archive_endpoint +from services.bots_orchestrator import BotsOrchestrator + +BOT_NAME = "hummingbot-pmm-1" + + +class _StubOrchestrator(BotsOrchestrator): + """Real stop_and_archive_bot, but no Docker/MQTT/database side effects.""" + + def __init__(self): # noqa: D107 - deliberately skips BotsOrchestrator.__init__ + self.active_bots = {BOT_NAME: {"status": "running"}} + self.stopping_bots = set() + self.mqtt_manager = SimpleNamespace(clear_bot_data=lambda name: self.calls["clear_bot_data"].append(name)) + self.db_manager = SimpleNamespace(get_session_context=self._session) + self.calls = { + "get_bot_status": [], + "mark_bot_run_stopped": [], + "stop_bot": [], + "clear_bot_data": [], + } + + @asynccontextmanager + async def _session(self): + yield None + + def get_bot_status(self, bot_name): + self.calls["get_bot_status"].append(bot_name) + return {"performance": {}} + + async def mark_bot_run_stopped(self, bot_name, final_status=None): + self.calls["mark_bot_run_stopped"].append(bot_name) + + async def stop_bot(self, bot_name, **kwargs): + self.calls["stop_bot"].append(bot_name) + return {"success": True} + + +class _StubDocker: + def __init__(self): + self.stopped = [] + self.status_checked = [] + self.removed = [] + + def stop_container(self, container_name): + self.stopped.append(container_name) + + def get_container_status(self, container_name): + self.status_checked.append(container_name) + return {"state": {"status": "exited"}} + + def remove_container(self, container_name, force=True): + self.removed.append((container_name, force)) + return {"success": True} + + +class _StubArchiver: + def __init__(self): + self.archived = [] + + def archive_locally(self, bot_name, instance_dir): + self.archived.append((bot_name, instance_dir)) + + def archive_and_upload(self, bot_name, instance_dir, bucket_name=None): + self.archived.append((bot_name, instance_dir, bucket_name)) + + +@pytest.fixture +def no_sleep(monkeypatch): + """The workflow sleeps 15s for graceful shutdown; skip the wall-clock wait.""" + async def _instant(_seconds): + return None + + monkeypatch.setattr(orchestrator_module.asyncio, "sleep", _instant) + + +@pytest.fixture +def archived_runs(monkeypatch): + """Capture the name passed to the ARCHIVED bot-run update.""" + seen = [] + + class _StubRepo: + def __init__(self, session): + pass + + async def update_bot_run_archived(self, bot_name): + seen.append(bot_name) + + async def update_bot_run_stopped(self, bot_name, error_message=None): + seen.append(("stopped", bot_name, error_message)) + + monkeypatch.setattr(orchestrator_module, "BotRunRepository", _StubRepo) + return seen + + +async def test_every_step_uses_the_bot_name_verbatim(no_sleep, archived_runs): + """All 8 steps address the bot by the exact name they were given.""" + orch = _StubOrchestrator() + docker_manager = _StubDocker() + bot_archiver = _StubArchiver() + + await orch.stop_and_archive_bot( + bot_name=BOT_NAME, + skip_order_cancellation=True, + archive_locally=True, + s3_bucket=None, + docker_manager=docker_manager, + bot_archiver=bot_archiver, + ) + + # Steps 1-3: MQTT-side identity + assert orch.calls["get_bot_status"] == [BOT_NAME] + assert orch.calls["mark_bot_run_stopped"] == [BOT_NAME] + assert orch.calls["stop_bot"] == [BOT_NAME] + # Steps 5 and 7: container identity + assert docker_manager.stopped == [BOT_NAME] + assert docker_manager.status_checked == [BOT_NAME] + assert docker_manager.removed == [(BOT_NAME, False)] + # Step 6: archive identity and instance directory + assert bot_archiver.archived == [(BOT_NAME, os.path.join("bots", "instances", BOT_NAME))] + # Step 8: bot run marked archived + assert archived_runs == [BOT_NAME] + # Cleanup: stopping flag cleared and bot dropped from active_bots + assert orch.stopping_bots == set() + assert orch.calls["clear_bot_data"] == [BOT_NAME] + assert BOT_NAME not in orch.active_bots + + +async def test_s3_archive_uses_the_bot_name_verbatim(no_sleep, archived_runs): + """The S3 branch archives under the same unmodified name.""" + orch = _StubOrchestrator() + bot_archiver = _StubArchiver() + + await orch.stop_and_archive_bot( + bot_name=BOT_NAME, + skip_order_cancellation=True, + archive_locally=False, + s3_bucket="my-bucket", + docker_manager=_StubDocker(), + bot_archiver=bot_archiver, + ) + + assert bot_archiver.archived == [ + (BOT_NAME, os.path.join("bots", "instances", BOT_NAME), "my-bucket") + ] + + +async def test_endpoint_schedules_the_background_task_with_the_path_name(): + """The endpoint hands the background task the path parameter, unmodified.""" + background_tasks = BackgroundTasks() + bots_manager = SimpleNamespace( + active_bots={BOT_NAME: {}}, + stop_and_archive_bot="sentinel", + ) + docker_manager = object() + bot_archiver = object() + + response = await stop_and_archive_endpoint( + bot_name=BOT_NAME, + background_tasks=background_tasks, + bots_manager=bots_manager, + docker_manager=docker_manager, + bot_archiver=bot_archiver, + ) + + assert response["status"] == "success" + assert response["details"]["bot_name"] == BOT_NAME + assert len(background_tasks.tasks) == 1 + task = background_tasks.tasks[0] + assert task.func == "sentinel" + assert task.kwargs["bot_name"] == BOT_NAME + assert task.kwargs["docker_manager"] is docker_manager + assert task.kwargs["bot_archiver"] is bot_archiver + # The three aliases are gone: the name is passed exactly once. + assert [k for k in task.kwargs if "name" in k] == ["bot_name"] + + +async def test_endpoint_reports_not_found_for_an_inactive_bot(): + """An unknown bot is refused before anything destructive is scheduled.""" + background_tasks = BackgroundTasks() + bots_manager = SimpleNamespace(active_bots={"other-bot": {}}, stop_and_archive_bot="sentinel") + + response = await stop_and_archive_endpoint( + bot_name=BOT_NAME, + background_tasks=background_tasks, + bots_manager=bots_manager, + docker_manager=object(), + bot_archiver=object(), + ) + + assert response["status"] == "error" + assert response["details"]["bot_name"] == BOT_NAME + assert response["details"]["active_bots"] == ["other-bot"] + assert background_tasks.tasks == [] diff --git a/test/test_stop_and_archive_failure_paths.py b/test/test_stop_and_archive_failure_paths.py new file mode 100644 index 00000000..85e080b2 --- /dev/null +++ b/test/test_stop_and_archive_failure_paths.py @@ -0,0 +1,176 @@ +""" +Tests for the stop-and-archive failure paths. + +``BotsOrchestrator.stop_and_archive_bot`` used to ``return`` silently when the +bot process failed to stop, or when the container was still alive after every +retry. The run row was left stopped-but-not-archived with no ``error_message``, +and the caller had already been told the background task started fine. These +tests pin that both paths now close the run out with an error state. + +Run with: pytest test/test_stop_and_archive_failure_paths.py -v +""" +from contextlib import asynccontextmanager +from types import SimpleNamespace + +import pytest + +import services.bots_orchestrator as orchestrator_module +from services.bots_orchestrator import BotsOrchestrator + +BOT_NAME = "hummingbot-pmm-1" + + +class _StubOrchestrator(BotsOrchestrator): + """Real stop_and_archive_bot, but no Docker/MQTT/database side effects.""" + + def __init__(self, stop_response): # noqa: D107 - deliberately skips BotsOrchestrator.__init__ + self.active_bots = {BOT_NAME: {"status": "running"}} + self.stopping_bots = set() + self.mqtt_manager = SimpleNamespace(clear_bot_data=lambda name: None) + self.db_manager = SimpleNamespace(get_session_context=self._session) + self._stop_response = stop_response + + @asynccontextmanager + async def _session(self): + yield None + + def get_bot_status(self, bot_name): + return {"performance": {}} + + async def mark_bot_run_stopped(self, bot_name, final_status=None): + return None + + async def stop_bot(self, bot_name, **kwargs): + return self._stop_response + + +class _StubDocker: + """Container that never reaches the ``exited`` state.""" + + def __init__(self): + self.stop_attempts = 0 + self.removed = [] + + def stop_container(self, container_name): + self.stop_attempts += 1 + + def get_container_status(self, container_name): + return {"state": {"status": "running"}} + + def remove_container(self, container_name, force=True): + self.removed.append((container_name, force)) + return {"success": True} + + +class _StubArchiver: + def __init__(self): + self.archived = [] + + def archive_locally(self, bot_name, instance_dir): + self.archived.append((bot_name, instance_dir)) + + def archive_and_upload(self, bot_name, instance_dir, bucket_name=None): + self.archived.append((bot_name, instance_dir, bucket_name)) + + +@pytest.fixture +def no_sleep(monkeypatch): + """The workflow sleeps between steps; skip the wall-clock waits.""" + async def _instant(_seconds): + return None + + monkeypatch.setattr(orchestrator_module.asyncio, "sleep", _instant) + + +@pytest.fixture +def run_updates(monkeypatch): + """Capture every terminal update applied to the bot run row.""" + seen = [] + + class _StubRepo: + def __init__(self, session): + pass + + async def update_bot_run_archived(self, bot_name): + seen.append(("archived", bot_name)) + + async def update_bot_run_stopped(self, bot_name, final_status=None, error_message=None): + seen.append(("stopped", bot_name, error_message)) + + monkeypatch.setattr(orchestrator_module, "BotRunRepository", _StubRepo) + return seen + + +async def _run(orch, docker_manager): + await orch.stop_and_archive_bot( + bot_name=BOT_NAME, + skip_order_cancellation=True, + archive_locally=True, + s3_bucket=None, + docker_manager=docker_manager, + bot_archiver=_StubArchiver(), + ) + + +async def test_failed_bot_stop_marks_the_run_errored(no_sleep, run_updates): + """A refused bot stop leaves an error state naming the failed step.""" + orch = _StubOrchestrator({"success": False, "error": "mqtt timeout"}) + docker_manager = _StubDocker() + + await _run(orch, docker_manager) + + assert len(run_updates) == 1 + kind, bot_name, error_message = run_updates[0] + assert (kind, bot_name) == ("stopped", BOT_NAME) + assert "Failed to stop bot process" in error_message + assert "mqtt timeout" in error_message + # It really did bail out before touching the container. + assert docker_manager.stop_attempts == 0 + assert docker_manager.removed == [] + # Cleanup still runs. + assert orch.stopping_bots == set() + assert BOT_NAME not in orch.active_bots + + +async def test_no_stop_response_marks_the_run_errored(no_sleep, run_updates): + """A missing stop response is recorded too, not swallowed.""" + orch = _StubOrchestrator(None) + + await _run(orch, _StubDocker()) + + assert len(run_updates) == 1 + assert run_updates[0][0:2] == ("stopped", BOT_NAME) + assert "No response from bot orchestrator" in run_updates[0][2] + + +async def test_container_stop_retry_exhaustion_marks_the_run_errored(no_sleep, run_updates): + """A container that never exits leaves an error naming the retry count.""" + orch = _StubOrchestrator({"success": True}) + docker_manager = _StubDocker() + + await _run(orch, docker_manager) + + assert docker_manager.stop_attempts == 10 + assert len(run_updates) == 1 + kind, bot_name, error_message = run_updates[0] + assert (kind, bot_name) == ("stopped", BOT_NAME) + assert error_message == "Failed to stop container after 10 attempts" + # It bailed out before archiving or removing anything. + assert docker_manager.removed == [] + assert orch.stopping_bots == set() + assert BOT_NAME not in orch.active_bots + + +async def test_database_failure_while_recording_the_error_is_swallowed(no_sleep, monkeypatch): + """Recording the error must never raise out of the background task.""" + class _ExplodingRepo: + def __init__(self, session): + raise RuntimeError("database is down") + + monkeypatch.setattr(orchestrator_module, "BotRunRepository", _ExplodingRepo) + orch = _StubOrchestrator({"success": False, "error": "mqtt timeout"}) + + await _run(orch, _StubDocker()) + + assert orch.stopping_bots == set() + assert BOT_NAME not in orch.active_bots diff --git a/test/test_ticker_sources.py b/test/test_ticker_sources.py new file mode 100644 index 00000000..0ea0f601 --- /dev/null +++ b/test/test_ticker_sources.py @@ -0,0 +1,668 @@ +""" +Tests for the ticker source adapters and the on-demand ticker fetch. + +The volume-unit assertions here are regressions: `ascend_ex` and `okx_perpetual` both report +BASE volume in a field that used to be stored as quote volume, which made cross-exchange +liquidity comparison meaningless. + +Run with: pytest test/test_ticker_sources.py -v +""" +from decimal import Decimal +from typing import Any, Dict + +import pytest +from hummingbot.connector.exchange_base import ExchangeBase + +from services.ticker_sources import ( + TICKER_SPECS, + Ticker, + TickerFetchError, + TickerUnsupportedError, + _generic, + _normalize, + _request, + _spec_extract, + fetch_tickers, +) + + +class FakeConnector: + """Minimal stand-in for a hummingbot connector: canned payloads plus a symbol map.""" + + def __init__(self, symbol_map: Dict[str, str], get_payload=None, post_payload=None): + self._symbol_map = symbol_map + self._get_payload = get_payload + self._post_payload = post_payload + self.get_calls = [] + self.post_calls = [] + + async def trading_pair_symbol_map(self): + return self._symbol_map + + async def _api_get(self, **kwargs): + self.get_calls.append(kwargs) + payload = self._get_payload + return payload(kwargs) if callable(payload) else payload + + async def _api_post(self, **kwargs): + self.post_calls.append(kwargs) + payload = self._post_payload + return payload(kwargs) if callable(payload) else payload + + +async def run_spec(connector_name: str, connector: FakeConnector) -> Dict[str, Ticker]: + """Drive one spec end to end, the way _fetch does.""" + spec = TICKER_SPECS[connector_name] + rows = await _request(connector, spec) + return _normalize(connector.symbol_map_for_test, rows, _spec_extract(spec), connector_name) + + +# Attach the map used by run_spec without threading it through every call site. +FakeConnector.symbol_map_for_test = property(lambda self: self._symbol_map) + + +# ==================== Ticker volume derivation ==================== + +def test_quote_volume_derived_from_base(): + t = Ticker(price=Decimal("100"), base_volume=Decimal("5"), timestamp=1.0) + assert t.quote_volume == Decimal("500") + + +def test_base_volume_derived_from_quote(): + t = Ticker(price=Decimal("100"), quote_volume=Decimal("500"), timestamp=1.0) + assert t.base_volume == Decimal("5") + + +def test_reported_volumes_are_never_overwritten(): + t = Ticker( + price=Decimal("100"), base_volume=Decimal("5"), quote_volume=Decimal("999"), timestamp=1.0 + ) + assert t.base_volume == Decimal("5") + assert t.quote_volume == Decimal("999") + + +# ==================== Spec-driven adapters ==================== + +@pytest.mark.asyncio +async def test_binance_reports_both_volumes(): + connector = FakeConnector( + {"BTCUSDT": "BTC-USDT"}, + get_payload=[{ + "symbol": "BTCUSDT", "bidPrice": "100.0", "askPrice": "102.0", + "lastPrice": "101.5", "volume": "10", "quoteVolume": "1010", + }], + ) + tickers = await run_spec("binance", connector) + ticker = tickers["BTC-USDT"] + assert ticker.price == Decimal("101") # mid of bid/ask, not lastPrice + assert ticker.base_volume == Decimal("10") + assert ticker.quote_volume == Decimal("1010") + + +@pytest.mark.asyncio +async def test_binance_perpetual_falls_back_to_last_price(): + # The futures 24hr ticker carries no bid/ask. + connector = FakeConnector( + {"BTCUSDT": "BTC-USDT"}, + get_payload=[{ + "symbol": "BTCUSDT", "lastPrice": "63457.9", + "volume": "162487.091", "quoteVolume": "10447087916.01", + }], + ) + tickers = await run_spec("binance_perpetual", connector) + assert tickers["BTC-USDT"].price == Decimal("63457.9") + assert tickers["BTC-USDT"].quote_volume == Decimal("10447087916.01") + + +@pytest.mark.asyncio +async def test_ascend_ex_volume_is_base_not_quote(): + """Regression: `volume` is BASE volume; it used to be stored as the quote volume.""" + connector = FakeConnector( + {"BTC/USDT": "BTC-USDT"}, + get_payload={"data": [{ + "symbol": "BTC/USDT", "bid": ["100.0", "3"], "ask": ["102.0", "4"], "volume": "10", + }]}, + ) + tickers = await run_spec("ascend_ex", connector) + ticker = tickers["BTC-USDT"] + assert ticker.base_volume == Decimal("10") + assert ticker.quote_volume == Decimal("1010") # 10 * mid(101) + + +@pytest.mark.asyncio +async def test_okx_perpetual_volccy_is_base_and_vol24h_is_ignored(): + """Regression: for instType=SWAP, volCcy24h is BASE volume and vol24h counts contracts.""" + connector = FakeConnector( + {"BTC-USDT-SWAP": "BTC-USDT"}, + get_payload={"data": [{ + "instId": "BTC-USDT-SWAP", "bidPx": "100.0", "askPx": "102.0", "last": "101.5", + "vol24h": "9999999", "volCcy24h": "10", + }]}, + ) + tickers = await run_spec("okx_perpetual", connector) + ticker = tickers["BTC-USDT"] + assert ticker.base_volume == Decimal("10") # volCcy24h, never vol24h + assert ticker.quote_volume == Decimal("1010") + + +@pytest.mark.asyncio +async def test_okx_spot_volccy_is_quote_volume(): + """The same field means quote volume on SPOT, which is why the two specs differ.""" + connector = FakeConnector( + {"BTC-USDT": "BTC-USDT"}, + get_payload={"data": [{ + "instId": "BTC-USDT", "bidPx": "100.0", "askPx": "102.0", "last": "101.5", + "vol24h": "10", "volCcy24h": "1010", + }]}, + ) + ticker = (await run_spec("okx", connector))["BTC-USDT"] + assert ticker.base_volume == Decimal("10") + assert ticker.quote_volume == Decimal("1010") + + +@pytest.mark.asyncio +async def test_bybit_perpetual_passes_dict_path_and_no_limit_id(): + """bybit_perpetual's _api_request indexes the endpoint by market, so it needs the dict.""" + connector = FakeConnector( + {"BTCUSDT": "BTC-USDT"}, + get_payload={"result": {"list": [{ + "symbol": "BTCUSDT", "bid1Price": "100.0", "ask1Price": "102.0", + "lastPrice": "101.5", "volume24h": "10", "turnover24h": "1010", + }]}}, + ) + await run_spec("bybit_perpetual", connector) + call = connector.get_calls[0] + assert isinstance(call["path_url"], dict) + assert call["params"] == {"category": "linear"} + # The connector computes its own throttler id when limit_id is absent. + assert "limit_id" not in call + + +@pytest.mark.asyncio +async def test_kraken_keyed_dict_and_24h_base_volume(): + # `v` and `p` are [today, last_24h] pairs, so the 24h figure is index 1. + connector = FakeConnector( + {"XBTUSDT": "BTC-USDT"}, + get_payload={"XBTUSDT": { + "a": ["102.0", "1", "1.0"], "b": ["100.0", "1", "1.0"], "c": ["101.0", "0.5"], + "v": ["3.0", "10.0"], + }}, + ) + ticker = (await run_spec("kraken", connector))["BTC-USDT"] + assert ticker.price == Decimal("101") + assert ticker.base_volume == Decimal("10.0") + assert ticker.quote_volume == Decimal("1010.0") + + +# ==================== Hyperliquid ==================== + +@pytest.mark.asyncio +async def test_hyperliquid_spot_keys_on_coin_not_index(): + """assetCtxs is longer than universe, so the two must not be zipped.""" + payload = [ + {"universe": [{"name": "PURR/USDC"}]}, + [ + {"coin": "@1", "midPx": "2.0", "dayBaseVlm": "5", "dayNtlVlm": "10"}, + {"coin": "PURR/USDC", "midPx": "0.065", "dayBaseVlm": "100", "dayNtlVlm": "6.5"}, + ], + ] + connector = FakeConnector( + {"PURR/USDC": "PURR-USDC", "@1": "UBTC-USDC"}, post_payload=payload + ) + tickers = await fetch_tickers(connector, "hyperliquid", raise_on_error=True) + assert set(tickers) == {"PURR-USDC", "UBTC-USDC"} + assert tickers["PURR-USDC"].price == Decimal("0.065") + assert tickers["PURR-USDC"].base_volume == Decimal("100") + assert tickers["PURR-USDC"].quote_volume == Decimal("6.5") + assert connector.post_calls[0]["data"] == {"type": "spotMetaAndAssetCtxs"} + + +@pytest.mark.asyncio +async def test_hyperliquid_perpetual_zips_universe_with_ctxs(): + payload = [ + {"universe": [{"name": "BTC"}, {"name": "ETH"}]}, + [ + {"midPx": "63460.5", "dayBaseVlm": "35116.8", "dayNtlVlm": "2259842755.8"}, + {"midPx": "3000.0", "dayBaseVlm": "1000", "dayNtlVlm": "3000000"}, + ], + ] + connector = FakeConnector({"BTC": "BTC-USD", "ETH": "ETH-USD"}, post_payload=payload) + tickers = await fetch_tickers(connector, "hyperliquid_perpetual", raise_on_error=True) + assert tickers["BTC-USD"].price == Decimal("63460.5") + assert tickers["BTC-USD"].quote_volume == Decimal("2259842755.8") + assert tickers["ETH-USD"].price == Decimal("3000.0") + assert connector.post_calls[0]["data"] == {"type": "metaAndAssetCtxs"} + + +class Hip3Connector(FakeConnector): + """Hyperliquid perp connector exposing the HIP-3 dex-market hooks.""" + + def __init__(self, symbol_map, post_payload, hip3_rows, fail=False): + super().__init__(symbol_map, post_payload=post_payload) + self._hip3_rows = hip3_rows + self._dex_markets = None # force a fetch rather than reusing the symbol-map cache + self._fail = fail + self.hip3_fetches = 0 + + async def _fetch_and_cache_hip3_market_data(self): + self.hip3_fetches += 1 + if self._fail: + raise RuntimeError("allPerpMetas unavailable") + return [{"name": "xyz"}] + + def _iter_hip3_merged_markets(self, dex_markets=None): + return iter(self._hip3_rows) + + +PERP_PAYLOAD = [ + {"universe": [{"name": "BTC"}]}, + [{"midPx": "63460.5", "dayBaseVlm": "35116.8", "dayNtlVlm": "2259842755.8"}], +] +HIP3_ROWS = [{"name": "xyz:TSLA", "midPx": "400.0", "dayBaseVlm": "10", "dayNtlVlm": "4000"}] +HIP3_MAP = {"BTC": "BTC-USD", "xyz:TSLA": "XYZ:TSLA-USD"} + + +@pytest.fixture(autouse=True) +def _clear_hip3_snapshots(): + from services import ticker_sources + + ticker_sources._hip3_snapshots.clear() + yield + ticker_sources._hip3_snapshots.clear() + + +@pytest.mark.asyncio +async def test_hip3_markets_are_included(): + """HIP-3 builder dexes are absent from metaAndAssetCtxs and were being dropped entirely.""" + connector = Hip3Connector(HIP3_MAP, PERP_PAYLOAD, HIP3_ROWS) + tickers = await fetch_tickers(connector, "hyperliquid_perpetual", raise_on_error=True) + assert set(tickers) == {"BTC-USD", "XYZ:TSLA-USD"} + assert tickers["XYZ:TSLA-USD"].price == Decimal("400.0") + assert tickers["XYZ:TSLA-USD"].base_volume == Decimal("10") + assert tickers["XYZ:TSLA-USD"].quote_volume == Decimal("4000") + + +@pytest.mark.asyncio +async def test_hip3_snapshot_is_reused_within_the_interval(): + """~10 requests per refresh is too expensive to repeat on every cycle.""" + connector = Hip3Connector(HIP3_MAP, PERP_PAYLOAD, HIP3_ROWS) + for _ in range(3): + tickers = await fetch_tickers(connector, "hyperliquid_perpetual", raise_on_error=True) + assert "XYZ:TSLA-USD" in tickers + assert connector.hip3_fetches == 1 + + +@pytest.mark.asyncio +async def test_hip3_snapshot_refreshes_after_the_interval(): + from services import ticker_sources + + connector = Hip3Connector(HIP3_MAP, PERP_PAYLOAD, HIP3_ROWS) + await fetch_tickers(connector, "hyperliquid_perpetual", raise_on_error=True) + _, rows = ticker_sources._hip3_snapshots["hyperliquid_perpetual"] + ticker_sources._hip3_snapshots["hyperliquid_perpetual"] = (0.0, rows) # expire it + await fetch_tickers(connector, "hyperliquid_perpetual", raise_on_error=True) + assert connector.hip3_fetches == 2 + + +@pytest.mark.asyncio +async def test_hip3_failure_leaves_main_dex_intact(): + """HIP-3 is supplementary; losing it must not lose the 232 main-dex markets.""" + connector = Hip3Connector(HIP3_MAP, PERP_PAYLOAD, HIP3_ROWS, fail=True) + tickers = await fetch_tickers(connector, "hyperliquid_perpetual", raise_on_error=True) + assert set(tickers) == {"BTC-USD"} + + +@pytest.mark.asyncio +async def test_hip3_can_be_disabled(monkeypatch): + from config import settings + + monkeypatch.setattr(settings.market_data, "hyperliquid_hip3_interval", 0) + connector = Hip3Connector(HIP3_MAP, PERP_PAYLOAD, HIP3_ROWS) + tickers = await fetch_tickers(connector, "hyperliquid_perpetual", raise_on_error=True) + assert set(tickers) == {"BTC-USD"} + assert connector.hip3_fetches == 0 + + +@pytest.mark.asyncio +async def test_hyperliquid_perpetual_refuses_misaligned_payload(): + """A length mismatch would attach every price to the wrong market.""" + payload = [{"universe": [{"name": "BTC"}, {"name": "ETH"}]}, [{"midPx": "1"}]] + connector = FakeConnector({"BTC": "BTC-USD"}, post_payload=payload) + with pytest.raises(TickerFetchError, match="length mismatch"): + await fetch_tickers(connector, "hyperliquid_perpetual", raise_on_error=True) + + +# ==================== Error paths ==================== + +@pytest.mark.asyncio +async def test_empty_symbol_map_raises_with_a_usable_message(): + connector = FakeConnector({}, get_payload=[]) + with pytest.raises(TickerFetchError, match="symbol map is empty"): + await fetch_tickers(connector, "binance", raise_on_error=True) + + +@pytest.mark.asyncio +async def test_rows_present_but_none_mapped_raises_with_samples(): + connector = FakeConnector( + {"BTCUSDT": "BTC-USDT"}, + get_payload=[{"symbol": "NOPE", "lastPrice": "1", "quoteVolume": "1"}], + ) + with pytest.raises(TickerFetchError) as excinfo: + await fetch_tickers(connector, "binance", raise_on_error=True) + message = str(excinfo.value) + assert "NOPE" in message and "BTCUSDT" in message + + +@pytest.mark.asyncio +async def test_unsupported_connector_is_rejected_before_any_request(): + connector = FakeConnector({"X": "X-USD"}) + with pytest.raises(TickerUnsupportedError, match="requires API credentials"): + await fetch_tickers(connector, "coinbase_advanced_trade", raise_on_error=True) + assert connector.get_calls == [] and connector.post_calls == [] + + +@pytest.mark.asyncio +async def test_background_mode_swallows_errors(): + connector = FakeConnector({}, get_payload=[]) + assert await fetch_tickers(connector, "binance") == {} + + +# ==================== Generic fallback ==================== + +class NoBulkConnector(FakeConnector): + """A connector using the base one-request-per-pair get_last_traded_prices.""" + + # Binding the base implementation makes _has_bulk_last_traded_prices report False. + get_last_traded_prices = ExchangeBase.get_last_traded_prices + + async def get_all_pairs_prices(self): + raise NotImplementedError + + +@pytest.mark.asyncio +async def test_generic_refuses_to_fan_out_one_request_per_pair(): + symbol_map = {f"SYM{i}": f"SYM{i}-USDT" for i in range(500)} + connector = NoBulkConnector(symbol_map) + with pytest.raises(TickerUnsupportedError, match="one HTTP request each"): + await _generic(connector, "some_exchange", symbol_map) + + +@pytest.mark.asyncio +async def test_generic_uses_symbol_map_not_trading_rules(): + """trading_rules stays empty forever on a keyless connector, so pairs come from the map.""" + calls = {} + + class BulkConnector(FakeConnector): + trading_rules: Dict[str, Any] = {} + + async def get_all_pairs_prices(self): + raise NotImplementedError + + async def get_last_traded_prices(self, pairs): + calls["pairs"] = pairs + return {pair: "10" for pair in pairs} + + symbol_map = {"BTCUSDT": "BTC-USDT", "ETHUSDT": "ETH-USDT"} + tickers = await _generic(BulkConnector(symbol_map), "some_exchange", symbol_map) + assert sorted(calls["pairs"]) == ["BTC-USDT", "ETH-USDT"] + assert tickers["BTC-USDT"].price == Decimal("10") + assert tickers["BTC-USDT"].quote_volume is None # never guessed on the generic path + + +# ==================== On-demand fetching in MarketDataService ==================== + +class StubConnectorService: + """Enough of UnifiedConnectorService for the on-demand ticker path.""" + + def __init__(self, known=("bybit",)): + self._known = set(known) + self._data_connectors: Dict[str, Any] = {} + + def is_known_connector(self, connector_name): + return connector_name in self._known + + def get_best_connector_for_market(self, connector_name, account_name=None): + self._data_connectors.setdefault(connector_name, object()) + return self._data_connectors[connector_name] + + def get_all_trading_connectors(self): + return {} + + +def make_service(monkeypatch, fetch_impl, known=("bybit",)): + from services import market_data_service as mds + + monkeypatch.setattr(mds, "fetch_tickers", fetch_impl) + return mds.MarketDataService(connector_service=StubConnectorService(known)) + + +@pytest.mark.asyncio +async def test_concurrent_requests_trigger_a_single_upstream_fetch(monkeypatch): + import asyncio + + calls = [] + + async def slow_fetch(connector, connector_name, *, raise_on_error=False): + calls.append(connector_name) + await asyncio.sleep(0.05) + return {"BTC-USDT": Ticker(price=Decimal("100"), timestamp=1.0)} + + service = make_service(monkeypatch, slow_fetch) + results = await asyncio.gather( + *[service.fetch_connector_tickers("bybit") for _ in range(10)] + ) + assert len(calls) == 1 + assert all(r["BTC-USDT"].price == Decimal("100") for r in results) + + +@pytest.mark.asyncio +async def test_fresh_cache_is_served_and_force_refetches(monkeypatch): + calls = [] + + async def fetch(connector, connector_name, *, raise_on_error=False): + calls.append(connector_name) + return {"BTC-USDT": Ticker(price=Decimal("100"), timestamp=1.0)} + + service = make_service(monkeypatch, fetch) + await service.fetch_connector_tickers("bybit") + await service.fetch_connector_tickers("bybit") # served from cache + assert len(calls) == 1 + + await service.fetch_connector_tickers("bybit", force=True) + assert len(calls) == 2 + + await service.fetch_connector_tickers("bybit", max_age=0) # stale immediately + assert len(calls) == 3 + + +@pytest.mark.asyncio +async def test_unknown_connector_raises(monkeypatch): + from services.unified_connector_service import UnknownConnectorError + + async def fetch(connector, connector_name, *, raise_on_error=False): + raise AssertionError("must not reach the connector") + + service = make_service(monkeypatch, fetch) + with pytest.raises(UnknownConnectorError): + await service.fetch_connector_tickers("not_a_real_exchange") + + +def test_paper_trade_and_testnet_are_not_market_data_connectors(): + from services.market_data_service import is_market_data_connector + + assert is_market_data_connector("binance") is True + assert is_market_data_connector("hyperliquid_perpetual") is True + assert is_market_data_connector("kucoin_hft") is True + for excluded in ( + "binance_paper_trade", "hyperliquid_testnet", "bybit_perpetual_testnet", + "architect_perpetual_sandbox", + ): + assert is_market_data_connector(excluded) is False, excluded + + +@pytest.mark.asyncio +async def test_testnet_request_is_refused_and_never_enters_the_pool(monkeypatch): + """Serving a testnet would cache it and pull it into _rebuild_price_pool.""" + async def fetch(connector, connector_name, *, raise_on_error=False): + raise AssertionError("must not fetch a testnet") + + service = make_service(monkeypatch, fetch, known=("bybit_testnet",)) + with pytest.raises(TickerUnsupportedError, match="not real market data"): + await service.fetch_connector_tickers("bybit_testnet") + assert service.get_tickers() == {} + assert service.prices == {} + + +def test_collection_set_skips_paper_trade_and_testnet(monkeypatch): + from services import market_data_service as mds + + service = mds.MarketDataService(connector_service=StubConnectorService()) + service._connector_service._data_connectors = { + "binance": object(), "bybit_testnet": object(), + "binance_paper_trade": object(), "okx": object(), + } + assert sorted(service._connected_connector_names()) == ["binance", "okx"] + + +@pytest.mark.asyncio +async def test_connector_construction_failure_is_a_clean_unsupported_error(monkeypatch): + """A missing optional dependency must not surface as a raw 500.""" + async def fetch(connector, connector_name, *, raise_on_error=False): + raise AssertionError("must not be reached") + + service = make_service(monkeypatch, fetch) + + def boom(connector_name, account_name=None): + raise ModuleNotFoundError("No module named 'v4_proto'") + + monkeypatch.setattr(service._connector_service, "get_best_connector_for_market", boom) + with pytest.raises(TickerUnsupportedError, match="cannot be instantiated"): + await service.fetch_connector_tickers("bybit") + + +@pytest.mark.asyncio +async def test_on_demand_connector_joins_background_collection(monkeypatch): + async def fetch(connector, connector_name, *, raise_on_error=False): + return {"BTC-USDT": Ticker(price=Decimal("100"), timestamp=1.0)} + + service = make_service(monkeypatch, fetch) + await service.fetch_connector_tickers("bybit") + assert "bybit" in service._connected_connector_names() + + +def test_is_more_liquid_prefers_known_and_larger_quote_volume(): + from services.market_data_service import MarketDataService as MDS + + high = Ticker(price=Decimal("1"), quote_volume=Decimal("100"), timestamp=1.0) + low = Ticker(price=Decimal("1"), quote_volume=Decimal("10"), timestamp=1.0) + unknown_old = Ticker(price=Decimal("1"), timestamp=1.0) + unknown_new = Ticker(price=Decimal("1"), timestamp=2.0) + + assert MDS._is_more_liquid(high, low) is True + assert MDS._is_more_liquid(low, high) is False + # A base-only exchange still participates instead of always ranking as zero. + assert MDS._is_more_liquid(low, unknown_old) is True + assert MDS._is_more_liquid(unknown_old, low) is False + assert MDS._is_more_liquid(unknown_new, unknown_old) is True + + +# ==================== Merged /market-data/tickers endpoint ==================== + +def test_connector_filter_parsing(): + from routers.market_data import _requested_connectors + + assert _requested_connectors(None) == [] + assert _requested_connectors(["binance,okx"]) == ["binance", "okx"] # comma-separated + assert _requested_connectors(["binance", "okx"]) == ["binance", "okx"] # repeated param + assert _requested_connectors([" binance , okx "]) == ["binance", "okx"] # whitespace + assert _requested_connectors(["binance,,okx", "binance"]) == ["binance", "okx"] # blanks + dupes + + +class StubMarketDataService: + """Stands in for MarketDataService at the router boundary.""" + + def __init__(self, pool=None, results=None, errors=None): + self._pool = pool or {} + self._results = results or {} + self._errors = errors or {} + self.fetch_calls = [] + + def get_tickers(self): + return self._pool + + def ticker_updated_at(self, connector_name): + return 123.0 + + def collected_connector_names(self): + return sorted(self._pool) + + async def fetch_tickers_for(self, names, *, max_age=None, force=False): + from services.unified_connector_service import UnknownConnectorError + + self.fetch_calls.append((tuple(names), force)) + unknown = [n for n in names if n not in self._results and n not in self._errors] + if unknown: + raise UnknownConnectorError(f"Connectors not found: {', '.join(unknown)}") + return ( + {n: self._results[n] for n in names if n in self._results}, + {n: self._errors[n] for n in names if n in self._errors}, + ) + + +def client_for(service): + from fastapi import FastAPI + from fastapi.testclient import TestClient + + from deps import get_market_data_service + from routers import market_data + + app = FastAPI() + app.include_router(market_data.router) + app.dependency_overrides[get_market_data_service] = lambda: service + return TestClient(app, raise_server_exceptions=False) + + +TICK = {"BTC-USDT": Ticker(price=Decimal("100"), quote_volume=Decimal("5"), timestamp=1.0)} + + +def test_no_filter_reads_the_pool_without_fetching(): + service = StubMarketDataService(pool={"binance": TICK, "okx": TICK}) + body = client_for(service).get("/market-data/tickers").json() + assert set(body["tickers"]) == {"binance", "okx"} + assert body["counts"] == {"binance": 1, "okx": 1} + assert service.fetch_calls == [] # a plain pool read makes no requests + + +def test_multiple_connectors_are_fetched_together(): + service = StubMarketDataService(results={"binance": TICK, "okx": TICK, "htx": TICK}) + body = client_for(service).get("/market-data/tickers?connectors=binance,okx").json() + assert set(body["tickers"]) == {"binance", "okx"} + assert service.fetch_calls == [(("binance", "okx"), False)] # one gathered call + + +def test_refresh_forces_a_fetch_of_the_whole_pool(): + service = StubMarketDataService(pool={"binance": TICK}, results={"binance": TICK}) + client_for(service).get("/market-data/tickers?refresh=true") + assert service.fetch_calls == [(("binance",), True)] + + +def test_partial_failure_still_returns_the_successful_connectors(): + service = StubMarketDataService( + results={"binance": TICK}, errors={"ascend_ex": TickerFetchError("symbol map is empty")} + ) + response = client_for(service).get("/market-data/tickers?connectors=binance,ascend_ex") + assert response.status_code == 200 + body = response.json() + assert body["counts"] == {"binance": 1} + assert "symbol map is empty" in body["errors"]["ascend_ex"] + + +def test_status_codes_when_nothing_can_be_served(): + unsupported = StubMarketDataService(errors={"xrpl": TickerUnsupportedError("needs a node pool")}) + assert client_for(unsupported).get("/market-data/tickers?connectors=xrpl").status_code == 400 + + failed = StubMarketDataService(errors={"ascend_ex": TickerFetchError("boom")}) + assert client_for(failed).get("/market-data/tickers?connectors=ascend_ex").status_code == 502 + + empty = StubMarketDataService() + assert client_for(empty).get("/market-data/tickers?connectors=nope").status_code == 404 diff --git a/test/test_websocket_auth_channels.py b/test/test_websocket_auth_channels.py new file mode 100644 index 00000000..9b7c450b --- /dev/null +++ b/test/test_websocket_auth_channels.py @@ -0,0 +1,154 @@ +""" +Tests for WebSocket handshake authentication (SEC-059). + +Credentials must travel in headers only — never in the query string, which uvicorn writes +to its access log for every handshake — and an unauthenticated peer must be refused at the +handshake instead of being accepted and then closed with 4001. + +Run with: pytest test/test_websocket_auth_channels.py -v +""" +import base64 + +import pytest +from fastapi import FastAPI +from fastapi.testclient import TestClient +from starlette.testclient import WebSocketDenialResponse + +from config import settings +from routers.websocket import AUTH_SUBPROTOCOL, router + +WS_PATHS = ["/ws/market-data", "/ws/executors"] + + +class _StubManager: + """Minimal stand-in for the real managers: the tests never get past the handshake.""" + + def generate_connection_id(self): + return "test-conn" + + async def handle_subscribe(self, *args, **kwargs): + pass + + async def handle_unsubscribe(self, *args, **kwargs): + pass + + def remove_connection(self, conn_id): + pass + + +@pytest.fixture +def client(): + app = FastAPI() + app.include_router(router) + app.state.websocket_manager = _StubManager() + app.state.executor_ws_manager = _StubManager() + return TestClient(app) + + +@pytest.fixture +def credentials(): + return settings.security.username, settings.security.password + + +def _basic_header(user: str, password: str) -> dict: + blob = base64.b64encode(f"{user}:{password}".encode()).decode() + return {"Authorization": f"Basic {blob}"} + + +def _subprotocols(user: str, password: str) -> list: + blob = base64.urlsafe_b64encode(f"{user}:{password}".encode()).decode().rstrip("=") + return [AUTH_SUBPROTOCOL, blob] + + +class TestQueryParamCredentialsAreRejected: + """The query-string channel is gone; it must not authenticate anything.""" + + @pytest.mark.parametrize("path", WS_PATHS) + def test_username_password_query_params_are_rejected(self, client, credentials, path): + user, password = credentials + with pytest.raises(WebSocketDenialResponse) as exc: + with client.websocket_connect(f"{path}?username={user}&password={password}"): + pass + assert exc.value.status_code == 401 + + @pytest.mark.parametrize("path", WS_PATHS) + def test_token_query_param_is_rejected(self, client, credentials, path): + user, password = credentials + token = base64.b64encode(f"{user}:{password}".encode()).decode() + with pytest.raises(WebSocketDenialResponse) as exc: + with client.websocket_connect(f"{path}?token={token}"): + pass + assert exc.value.status_code == 401 + + def test_no_code_path_reads_credentials_from_the_query_string(self): + import inspect + + import routers.websocket as ws_module + + source = inspect.getsource(ws_module) + assert "query_params" not in source + + +class TestHeaderChannelsAuthenticate: + @pytest.mark.parametrize("path", WS_PATHS) + def test_authorization_basic_header_succeeds(self, client, credentials, path): + user, password = credentials + with client.websocket_connect(path, headers=_basic_header(user, password)) as ws: + message = ws.receive_json() + assert message["type"] == "connected" + + @pytest.mark.parametrize("path", WS_PATHS) + def test_browser_subprotocol_channel_succeeds_and_is_echoed(self, client, credentials, path): + user, password = credentials + with client.websocket_connect(path, subprotocols=_subprotocols(user, password)) as ws: + message = ws.receive_json() + assert ws.accepted_subprotocol == AUTH_SUBPROTOCOL + assert message["type"] == "connected" + + @pytest.mark.parametrize("path", WS_PATHS) + def test_padded_standard_base64_in_the_subprotocol_also_decodes(self, client, credentials, path): + user, password = credentials + blob = base64.b64encode(f"{user}:{password}".encode()).decode() + with client.websocket_connect(path, subprotocols=[AUTH_SUBPROTOCOL, blob]) as ws: + message = ws.receive_json() + assert message["type"] == "connected" + + +class TestUnauthenticatedIsRefusedAtTheHandshake: + @pytest.mark.parametrize("path", WS_PATHS) + def test_no_credentials_at_all(self, client, path): + with pytest.raises(WebSocketDenialResponse) as exc: + with client.websocket_connect(path): + pass + assert exc.value.status_code == 401 + assert "basic" in exc.value.headers.get("www-authenticate", "").lower() + + @pytest.mark.parametrize("path", WS_PATHS) + def test_wrong_password_in_the_header(self, client, credentials, path): + user, _ = credentials + with pytest.raises(WebSocketDenialResponse) as exc: + with client.websocket_connect(path, headers=_basic_header(user, "not-the-password")): + pass + assert exc.value.status_code == 401 + + @pytest.mark.parametrize("path", WS_PATHS) + def test_wrong_password_in_the_subprotocol(self, client, credentials, path): + user, _ = credentials + with pytest.raises(WebSocketDenialResponse) as exc: + with client.websocket_connect(path, subprotocols=_subprotocols(user, "nope")): + pass + assert exc.value.status_code == 401 + + @pytest.mark.parametrize("path", WS_PATHS) + def test_garbage_credentials_do_not_raise_out_of_the_handler(self, client, path): + with pytest.raises(WebSocketDenialResponse) as exc: + with client.websocket_connect(path, headers={"Authorization": "Basic !!!not-base64"}): + pass + assert exc.value.status_code == 401 + + @pytest.mark.parametrize("path", WS_PATHS) + def test_subprotocol_without_a_credential_blob(self, client, path): + with pytest.raises(WebSocketDenialResponse) as exc: + with client.websocket_connect(path, subprotocols=[AUTH_SUBPROTOCOL]): + pass + assert exc.value.status_code == 401 diff --git a/utils/hummingbot_database_reader.py b/utils/hummingbot_database_reader.py index 110d057d..972aa7ac 100644 --- a/utils/hummingbot_database_reader.py +++ b/utils/hummingbot_database_reader.py @@ -1,13 +1,13 @@ -import os -import pandas as pd import json -from typing import List, Dict, Any +import os +from typing import Any, Dict, List +import pandas as pd from hummingbot.core.data_type.common import TradeType from hummingbot.strategy_v2.models.base import RunnableStatus from hummingbot.strategy_v2.models.executors import CloseType from hummingbot.strategy_v2.models.executors_info import ExecutorInfo -from sqlalchemy import create_engine, insert, text, MetaData, Table, Column, VARCHAR, INT, FLOAT, Integer, String, Float +from sqlalchemy import create_engine, text from sqlalchemy.orm import sessionmaker @@ -25,7 +25,7 @@ def __init__(self, db_path: str): def _get_table_status(table_loader): try: data = table_loader() - return "Correct" if len(data) > 0 else f"Error - No records matched" + return "Correct" if len(data) > 0 else "Error - No records matched" except Exception as e: return f"Error - {str(e)}" @@ -38,7 +38,8 @@ def status(self): controller_status = self._get_table_status(self.get_controllers_data) positions_status = self._get_table_status(self.get_positions) general_status = all(status == "Correct" for status in - [trade_fill_status, orders_status, order_status_status, executors_status, controller_status, positions_status]) + [trade_fill_status, orders_status, order_status_status, + executors_status, controller_status, positions_status]) status = {"db_name": self.db_name, "db_path": self.db_path, "trade_fill": trade_fill_status, @@ -51,10 +52,32 @@ def status(self): } return status + @staticmethod + def _as_numeric(df: pd.DataFrame, columns: List[str]) -> pd.DataFrame: + """Force the given columns to a numeric dtype, in place. + + A table with zero rows comes back from read_sql_query with `object` columns -- + pandas has nothing to infer a dtype from -- and object dtype is not merely + cosmetic downstream: `.cumsum()` raises `TypeError: cumsum is not supported for + object dtype` on it however empty the frame is. That is a real archived bot, not a + synthetic one: a bot whose order size sits below the exchange minimum has every + order rejected and archives with an empty TradeFill table, which turned + /archived-bots/{db}/performance and /summary into 500s. + + Coercing per column rather than with DataFrame.apply is deliberate: apply never + calls the function on a frame with no rows, so it leaves exactly the case this + exists for untouched. + """ + for column in columns: + if column in df.columns: + df[column] = pd.to_numeric(df[column], errors="coerce") + return df + def get_orders(self): with self.session_maker() as session: query = "SELECT * FROM 'Order'" orders = pd.read_sql_query(text(query), session.connection()) + self._as_numeric(orders, ["amount", "price"]) orders["amount"] = orders["amount"] / 1e6 orders["price"] = orders["price"] / 1e6 orders.rename(columns={"market": "connector_name", "symbol": "trading_pair"}, inplace=True) @@ -67,6 +90,7 @@ def get_trade_fills(self): query = "SELECT * FROM TradeFill" trade_fills = pd.read_sql_query(text(query), session.connection()) trade_fills.rename(columns={"market": "connector_name", "symbol": "trading_pair"}, inplace=True) + self._as_numeric(trade_fills, float_cols) trade_fills[float_cols] = trade_fills[float_cols] / 1e6 trade_fills["cum_fees_in_quote"] = trade_fills.groupby(groupers)["trade_fee_in_quote"].cumsum() trade_fills["trade_fee"] = trade_fills.groupby(groupers)["cum_fees_in_quote"].diff() @@ -96,116 +120,117 @@ def get_positions(self) -> pd.DataFrame: positions = pd.read_sql_query(text(query), session.connection()) # Convert decimal fields from stored format (divide by 1e6) decimal_cols = ["volume_traded_quote", "amount", "breakeven_price", "unrealized_pnl_quote", "cum_fees_quote"] + self._as_numeric(positions, decimal_cols) positions[decimal_cols] = positions[decimal_cols] / 1e6 return positions def calculate_trade_based_performance(self) -> pd.DataFrame: """ Calculate trade-based performance metrics using vectorized pandas operations. - + Returns: DataFrame with rolling performance metrics calculated per trading pair. """ # Get trade fills data trades = self.get_trade_fills() - + if len(trades) == 0: return pd.DataFrame() - + # Sort by timestamp to ensure proper rolling calculation trades = trades.sort_values(['trading_pair', 'connector_name', 'timestamp']).copy() - + # Create buy/sell indicator columns trades['is_buy'] = (trades['trade_type'].str.upper() == 'BUY').astype(int) trades['is_sell'] = (trades['trade_type'].str.upper() == 'SELL').astype(int) - + # Calculate buy and sell amounts and values vectorized trades['buy_amount'] = trades['amount'] * trades['is_buy'] trades['sell_amount'] = trades['amount'] * trades['is_sell'] trades['buy_value'] = trades['price'] * trades['amount'] * trades['is_buy'] trades['sell_value'] = trades['price'] * trades['amount'] * trades['is_sell'] - + # Group by trading_pair and connector_name for rolling calculations grouper = ['trading_pair', 'connector_name'] - + # Calculate cumulative volumes and values trades['buy_volume'] = trades.groupby(grouper)['buy_amount'].cumsum() trades['sell_volume'] = trades.groupby(grouper)['sell_amount'].cumsum() trades['buy_value_cum'] = trades.groupby(grouper)['buy_value'].cumsum() trades['sell_value_cum'] = trades.groupby(grouper)['sell_value'].cumsum() - + # Calculate average prices (avoid division by zero) trades['buy_avg_price'] = trades['buy_value_cum'] / trades['buy_volume'].replace(0, pd.NA) trades['sell_avg_price'] = trades['sell_value_cum'] / trades['sell_volume'].replace(0, pd.NA) - + # Forward fill average prices within each group to handle NaN values trades['buy_avg_price'] = trades.groupby(grouper)['buy_avg_price'].ffill().fillna(0) trades['sell_avg_price'] = trades.groupby(grouper)['sell_avg_price'].ffill().fillna(0) - + # Calculate net position trades['net_position'] = trades['buy_volume'] - trades['sell_volume'] - + # Calculate realized PnL trades['realized_trade_pnl_pct'] = ( (trades['sell_avg_price'] - trades['buy_avg_price']) / trades['buy_avg_price'] ).fillna(0) - + # Matched volume for realized PnL (minimum of buy and sell volumes) trades['matched_volume'] = pd.concat([trades['buy_volume'], trades['sell_volume']], axis=1).min(axis=1) trades['realized_trade_pnl_quote'] = trades['realized_trade_pnl_pct'] * trades['matched_volume'] * trades['buy_avg_price'] - + # Calculate unrealized PnL based on position direction # For long positions (net_position > 0): use current price vs buy_avg_price # For short positions (net_position < 0): use sell_avg_price vs current price trades['unrealized_trade_pnl_pct'] = 0.0 - + # Long positions long_mask = trades['net_position'] > 0 trades.loc[long_mask, 'unrealized_trade_pnl_pct'] = ( - (trades.loc[long_mask, 'price'] - trades.loc[long_mask, 'buy_avg_price']) / + (trades.loc[long_mask, 'price'] - trades.loc[long_mask, 'buy_avg_price']) / trades.loc[long_mask, 'buy_avg_price'] ).fillna(0) - - # Short positions + + # Short positions short_mask = trades['net_position'] < 0 trades.loc[short_mask, 'unrealized_trade_pnl_pct'] = ( - (trades.loc[short_mask, 'sell_avg_price'] - trades.loc[short_mask, 'price']) / + (trades.loc[short_mask, 'sell_avg_price'] - trades.loc[short_mask, 'price']) / trades.loc[short_mask, 'sell_avg_price'] ).fillna(0) - + # Calculate unrealized PnL in quote currency trades['unrealized_trade_pnl_quote'] = 0.0 - + # Long positions: use buy_avg_price as reference long_mask = trades['net_position'] > 0 trades.loc[long_mask, 'unrealized_trade_pnl_quote'] = ( - trades.loc[long_mask, 'unrealized_trade_pnl_pct'] * - trades.loc[long_mask, 'net_position'].abs() * + trades.loc[long_mask, 'unrealized_trade_pnl_pct'] * + trades.loc[long_mask, 'net_position'].abs() * trades.loc[long_mask, 'buy_avg_price'] ) - - # Short positions: use sell_avg_price as reference + + # Short positions: use sell_avg_price as reference short_mask = trades['net_position'] < 0 trades.loc[short_mask, 'unrealized_trade_pnl_quote'] = ( - trades.loc[short_mask, 'unrealized_trade_pnl_pct'] * - trades.loc[short_mask, 'net_position'].abs() * + trades.loc[short_mask, 'unrealized_trade_pnl_pct'] * + trades.loc[short_mask, 'net_position'].abs() * trades.loc[short_mask, 'sell_avg_price'] ) - + # Fees are already in trade_fee_in_quote column trades['fees_quote'] = trades['trade_fee_in_quote'] - + # Calculate net PnL trades['net_pnl_quote'] = ( - trades['realized_trade_pnl_quote'] + - trades['unrealized_trade_pnl_quote'] - + trades['realized_trade_pnl_quote'] + + trades['unrealized_trade_pnl_quote'] - trades['fees_quote'] ) - + # Calculate cumulative volume in quote currency trades['volume_quote'] = trades['price'] * trades['amount'] trades['cum_volume_quote'] = trades.groupby(grouper)['volume_quote'].cumsum() - + # Select and return relevant columns result_columns = [ 'timestamp', 'price', 'amount', 'trade_type', 'trading_pair', 'connector_name', @@ -214,9 +239,8 @@ def calculate_trade_based_performance(self) -> pd.DataFrame: 'unrealized_trade_pnl_pct', 'unrealized_trade_pnl_quote', 'fees_quote', 'net_pnl_quote', 'volume_quote', 'cum_volume_quote' ] - - return trades[result_columns].sort_values('timestamp') + return trades[result_columns].sort_values('timestamp') class PerformanceDataSource: @@ -237,7 +261,8 @@ def executors_df(self): executors["level_id"] = executors["config"].apply(lambda x: x.get("level_id")) executors["bep"] = executors["custom_info"].apply(lambda x: x["current_position_average_price"]) executors["order_ids"] = executors["custom_info"].apply(lambda x: x.get("order_ids")) - executors["close_price"] = executors["custom_info"].apply(lambda x: x.get("close_price", x["current_position_average_price"])) + executors["close_price"] = executors["custom_info"].apply( + lambda x: x.get("close_price", x["current_position_average_price"])) executors["sl"] = executors["config"].apply(lambda x: x.get("stop_loss")).fillna(0) executors["tp"] = executors["config"].apply(lambda x: x.get("take_profit")).fillna(0) executors["tl"] = executors["config"].apply(lambda x: x.get("time_limit")).fillna(0) @@ -307,4 +332,4 @@ def ensure_timestamp_in_seconds(timestamp: float) -> float: return timestamp_int else: raise ValueError( - "Timestamp is not in a recognized format. Must be in seconds, milliseconds, microseconds or nanoseconds.") \ No newline at end of file + "Timestamp is not in a recognized format. Must be in seconds, milliseconds, microseconds or nanoseconds.") diff --git a/utils/mqtt_manager.py b/utils/mqtt_manager.py index b4c5eaec..b6973e5b 100644 --- a/utils/mqtt_manager.py +++ b/utils/mqtt_manager.py @@ -2,7 +2,7 @@ import json import logging import time -from collections import defaultdict, deque +from collections import OrderedDict, defaultdict, deque from contextlib import asynccontextmanager from typing import Any, Callable, Dict, Optional, Set @@ -33,10 +33,12 @@ def __init__(self, host: str, port: int, username: str, password: str): # Auto-discovered bots self._discovered_bots: Dict[str, float] = {} # bot_id: last_seen_timestamp - - # Message deduplication tracking - self._processed_messages: Dict[str, float] = {} # message_hash: timestamp + + # Message deduplication tracking. Entries are inserted in non-decreasing + # timestamp order, so expiry only ever has to touch the oldest end. + self._processed_messages: "OrderedDict[str, float]" = OrderedDict() # message_hash: timestamp self._message_ttl = 300 # 5 minutes TTL for processed messages + self._max_processed_messages = 10000 # hard cap, so a log burst cannot grow the cache without bound # Connection state self._connected = False @@ -92,7 +94,7 @@ async def _get_client(self): for topic, qos in self._subscriptions: await client.subscribe(topic, qos=qos) yield client - + # Cleanup on exit self._connected = False @@ -207,14 +209,14 @@ async def _handle_log(self, bot_id: str, data: Any): level = data.get("level_name") or data.get("levelname") or data.get("level", "INFO") message = data.get("msg") or data.get("message", "") timestamp = data.get("timestamp") or data.get("time") or time.time() - + # Create hash for deduplication (bot_id + message + timestamp within 1 second) message_hash = f"{bot_id}:{message}:{int(timestamp)}" elif isinstance(data, str): message = data timestamp = time.time() level = "INFO" - + # Create hash for string messages message_hash = f"{bot_id}:{message}:{int(timestamp)}" else: @@ -227,13 +229,20 @@ async def _handle_log(self, bot_id: str, data: Any): logger.debug(f"Skipping duplicate log message from {bot_id}: {message[:50]}...") return - # Clean up old message hashes (older than TTL) - expired_hashes = [h for h, t in self._processed_messages.items() if current_time - t > self._message_ttl] - for h in expired_hashes: - del self._processed_messages[h] - - # Record this message as processed + # Clean up old message hashes (older than TTL). The cache is ordered by + # insertion time, so we only pop from the oldest end and stop at the + # first entry that is still live: O(expired) instead of O(cache size). + while self._processed_messages: + oldest_time = next(iter(self._processed_messages.values())) + if current_time - oldest_time <= self._message_ttl: + break + self._processed_messages.popitem(last=False) + + # Record this message as processed, evicting the oldest entries if the + # dedup cache has outgrown its cap self._processed_messages[message_hash] = current_time + while len(self._processed_messages) > self._max_processed_messages: + self._processed_messages.popitem(last=False) # Process the message if isinstance(data, dict): @@ -272,7 +281,7 @@ async def _handle_events(self, bot_id: str, data: Any): async def _handle_external_event(self, bot_id: str, channel: str, data: Any): """Handle external events.""" - event_type = channel.split("/")[-1] + # Process external events as needed async def _handle_rpc_response(self, topic: str, message): """Handle RPC responses on hummingbot-api/response/* topics.""" @@ -296,10 +305,8 @@ async def _handle_rpc_response(self, topic: str, message): async def _handle_command_response(self, bot_id: str, channel: str, data: Any): """Handle command responses (legacy - keeping for backward compatibility).""" - # Extract command from response channel (e.g., response/start/1234567890 or response/history) - channel_parts = channel.split("/") - if len(channel_parts) >= 2: - command = channel_parts[1] + # The command lives in the response channel (e.g. response/start/1234567890 + # or response/history); nothing consumes it yet. async def start(self): """Start the MQTT client."""