Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
160 changes: 160 additions & 0 deletions tests/tui_gateway/test_restore_session_usage.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,160 @@
"""Tests for _restore_session_usage — token counter restore on session.resume.

When the TUI gateway resumes a session, a fresh agent is built whose
session_*_tokens counters all start at zero. _restore_session_usage copies
the cumulative counts from the stored session row so that session.info /
session.usage reflects actual usage instead of 0/1.0M-0%.
"""

import pytest


class _FakeAgent:
"""Minimal stand-in with the same counter attributes as AIAgent."""

def __init__(self):
self.session_input_tokens = 0
self.session_output_tokens = 0
self.session_cache_read_tokens = 0
self.session_cache_write_tokens = 0
self.session_reasoning_tokens = 0
self.session_total_tokens = 0
self.session_api_calls = 0
self.session_estimated_cost_usd = 0.0
self.session_cost_status = "unknown"


def _import_restore():
from tui_gateway.server import _restore_session_usage
return _restore_session_usage


# ── Unit tests for _restore_session_usage ──────────────────────────


def test_restore_basic_token_counts():
"""All five token fields are populated from the stored session row."""
restore = _import_restore()
agent = _FakeAgent()
stored = {
"input_tokens": 12000,
"output_tokens": 3500,
"cache_read_tokens": 800,
"cache_write_tokens": 200,
"reasoning_tokens": 1500,
"api_call_count": 7,
}
restore(agent, stored)

assert agent.session_input_tokens == 12000
assert agent.session_output_tokens == 3500
assert agent.session_cache_read_tokens == 800
assert agent.session_cache_write_tokens == 200
assert agent.session_reasoning_tokens == 1500
assert agent.session_total_tokens == 12000 + 3500 + 800 + 200 + 1500
assert agent.session_api_calls == 7


def test_restore_total_is_sum():
"""session_total_tokens equals the sum of all five token components."""
restore = _import_restore()
agent = _FakeAgent()
stored = {
"input_tokens": 100,
"output_tokens": 200,
"cache_read_tokens": 300,
"cache_write_tokens": 400,
"reasoning_tokens": 500,
}
restore(agent, stored)
assert agent.session_total_tokens == 100 + 200 + 300 + 400 + 500


def test_restore_handles_none_values():
"""Missing or None fields default to zero."""
restore = _import_restore()
agent = _FakeAgent()
stored = {
"input_tokens": None,
"output_tokens": 50,
}
restore(agent, stored)

assert agent.session_input_tokens == 0
assert agent.session_output_tokens == 50
assert agent.session_cache_read_tokens == 0
assert agent.session_cache_write_tokens == 0
assert agent.session_reasoning_tokens == 0
assert agent.session_total_tokens == 50
assert agent.session_api_calls == 0


def test_restore_preserves_cost_estimate():
"""estimated_cost_usd and cost_status are carried over when present."""
restore = _import_restore()
agent = _FakeAgent()
stored = {
"input_tokens": 1000,
"output_tokens": 500,
"estimated_cost_usd": 0.042,
"cost_status": "estimated",
}
restore(agent, stored)

assert agent.session_estimated_cost_usd == pytest.approx(0.042)
assert agent.session_cost_status == "estimated"


def test_restore_skips_cost_when_absent():
"""Cost fields stay at defaults when not in the stored row."""
restore = _import_restore()
agent = _FakeAgent()
stored = {"input_tokens": 100, "output_tokens": 50}
restore(agent, stored)

assert agent.session_estimated_cost_usd == 0.0
assert agent.session_cost_status == "unknown"


def test_restore_empty_dict():
"""An empty stored dict leaves the agent at zero counters."""
restore = _import_restore()
agent = _FakeAgent()
restore(agent, {})

assert agent.session_input_tokens == 0
assert agent.session_output_tokens == 0
assert agent.session_total_tokens == 0
assert agent.session_api_calls == 0


# ── _get_usage picks up restored counters ──────────────────────────


def test_get_usage_reflects_restored_counters():
"""After restore, _get_usage returns the stored token counts."""
from tui_gateway.server import _get_usage

agent = _FakeAgent()
restore = _import_restore()
stored = {
"input_tokens": 5000,
"output_tokens": 2000,
"cache_read_tokens": 100,
"cache_write_tokens": 50,
"reasoning_tokens": 300,
"api_call_count": 3,
"estimated_cost_usd": 0.015,
"cost_status": "estimated",
}
restore(agent, stored)

usage = _get_usage(agent)

assert usage["input"] == 5000
assert usage["output"] == 2000
assert usage["cache_read"] == 100
assert usage["cache_write"] == 50
assert usage["reasoning"] == 300
assert usage["total"] == 5000 + 2000 + 100 + 50 + 300
assert usage["calls"] == 3
37 changes: 37 additions & 0 deletions tui_gateway/server.py
Original file line number Diff line number Diff line change
Expand Up @@ -1935,6 +1935,38 @@ def _get_usage(agent) -> dict:
return usage


def _restore_session_usage(agent, stored: dict) -> None:
"""Restore cumulative token counters from a stored session row.

When ``session.resume`` builds a fresh agent its counters start at zero.
The real cumulative usage lives in ``state.db`` (written by
``update_token_counts`` on every API call). Without this restore the
``session.info`` / ``session.usage`` payloads report all-zero usage
which shows as ``0/1.0M-0%`` in the Desktop status bar for gateway
sessions viewed after the fact.
"""
agent.session_input_tokens = stored.get("input_tokens") or 0
agent.session_output_tokens = stored.get("output_tokens") or 0
agent.session_cache_read_tokens = stored.get("cache_read_tokens") or 0
agent.session_cache_write_tokens = stored.get("cache_write_tokens") or 0
agent.session_reasoning_tokens = stored.get("reasoning_tokens") or 0
agent.session_total_tokens = (
agent.session_input_tokens
+ agent.session_output_tokens
+ agent.session_cache_read_tokens
+ agent.session_cache_write_tokens
+ agent.session_reasoning_tokens
)
agent.session_api_calls = stored.get("api_call_count") or 0
# Preserve cost estimate if available
cost_usd = stored.get("estimated_cost_usd")
if cost_usd:
agent.session_estimated_cost_usd = float(cost_usd)
cost_status = stored.get("cost_status")
if cost_status:
agent.session_cost_status = cost_status


def _probe_credentials(agent) -> str:
"""Light credential check at session creation — returns warning or ''."""
try:
Expand Down Expand Up @@ -3585,6 +3617,11 @@ def _(rid, params: dict) -> dict:
# state.db; home override is active here so config/skills/model
# resolve to the profile too.
agent = _make_agent(sid, target, session_id=target, session_db=db)
# Restore cumulative token counters from the stored session so
# that session.info / session.usage shows correct usage instead
# of all zeros when viewing gateway sessions in Desktop.

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Desktop no longer takes this eager branch: it omits eager_build in apps/desktop/src/app/session/hooks/use-session-actions/index.ts:503-510, so current main defers through tui_gateway/server.py:5721-5767. _start_agent_build then creates the agent and emits session.info without a restore (tui_gateway/server.py:1382,1445). Carry the stored row into that deferred record and restore there too.

if found:
_restore_session_usage(agent, found)
finally:
_clear_session_context(tokens)
except Exception as e:
Expand Down
Loading