refactor: remove dead shims, stale files, and test duplication (#243)

This commit is contained in:
Avanika Narayan
2026-04-13 15:03:03 -07:00
committed by GitHub
parent 171e78be5d
commit a3789212ec
65 changed files with 352 additions and 1230 deletions
Binary file not shown.

Before

Width:  |  Height:  |  Size: 235 KiB

-240
View File
@@ -1,240 +0,0 @@
# NeurIPS 2026 Experiment Plan: IPW/IPJ for Local AI
## Overview
Evaluate and optimize local AI models as OpenClaw agent brains, measuring
accuracy, latency, cost, energy, and FLOPs across 7 benchmarks.
## Results Storage
All results stored under `results/neurips-2026/`:
```
results/neurips-2026/
├── baselines/ # Step 1: Raw model scores
│ ├── {model}/{benchmark}/ # e.g. qwen-9b/pinchbench/
│ │ ├── results.jsonl # Per-task results
│ │ ├── summary.json # Aggregate metrics
│ │ └── telemetry.json # Energy, power, FLOPs, tokens
│ └── ...
├── agent-optimization/ # Step 2a: Agent improvements
│ ├── gepa/ # GEPA prompt evolution results
│ │ ├── generation_{N}/ # Per-generation best prompts
│ │ └── best_configs/ # Final optimized agent configs
│ ├── dspy/ # DSPy optimization results
│ │ ├── bootstrap/ # BootstrapFewShot results
│ │ └── mipro/ # MIPROv2 results
│ └── agent-configs/ # New agent configurations tested
├── intelligence-optimization/ # Step 2b: Model improvements
│ ├── sft/ # Supervised fine-tuning
│ │ ├── qwen-2b/ # Per-model training runs
│ │ ├── qwen-9b/
│ │ └── qwen-27b/
│ ├── lora/ # LoRA fine-tuning
│ │ ├── qwen-2b/
│ │ ├── qwen-9b/
│ │ └── qwen-27b/
│ └── rl/ # Reinforcement learning (GRPO)
│ ├── qwen-2b/
│ ├── qwen-9b/
│ └── qwen-27b/
├── optimized-eval/ # Step 3: Full eval with best configs
│ ├── {model}/{benchmark}/ # Same structure as baselines/
│ └── ...
└── analysis/ # Charts, tables, comparisons
├── pareto_frontier.json # IPW/IPJ data points
├── scaling_curves.json # Accuracy vs model size
├── cost_comparison.json # Local vs cloud economics
└── figures/ # Generated plots
```
## Hardware Stacks
Run eval metrics across three hardware vendor stacks to show
platform-agnostic IPW/IPJ results:
| Stack | Server-Class | Workstation/Consumer |
|-------|-------------|---------------------|
| **NVIDIA** | DGX Spark | RTX 6000 Pro |
| **AMD** | MI300x, MI355x | — |
| **Apple** | — | Mac Mini M4, Mac Studio M4 |
Results stored per-stack under each model's directory:
```
results/neurips-2026/baselines/{model}/{benchmark}/
├── nvidia-dgxspark/
│ ├── results.jsonl
│ ├── telemetry.json # NVML energy, GPU util, power
│ └── summary.json
├── nvidia-rtx6000pro/
├── amd-mi300x/
│ ├── telemetry.json # ROCm energy, GPU util, power
│ └── ...
├── amd-mi355x/
├── apple-macmini-m4/
│ ├── telemetry.json # Apple powermetrics energy
│ └── ...
└── apple-macstudio-m4/
```
OpenJarvis telemetry already supports all three vendors:
- NVIDIA: `telemetry/nvidia_monitor.py` (NVML)
- AMD: `telemetry/amd_monitor.py` (ROCm SMI)
- Apple: `telemetry/apple_monitor.py` (powermetrics)
Key comparisons:
- Same model, same benchmark, different hardware → IPW/IPJ per platform
- DGX Spark vs MI300x vs Mac Studio → server-class efficiency frontier
- RTX 6000 Pro vs Mac Mini M4 → consumer/workstation efficiency frontier
- GGUF models (Kimi, MiniMax) run on all platforms via llama.cpp/MLX
## Models (9 priority + 3 cloud baselines)
### Cloud Baselines
| ID | Model | Engine |
|----|-------|--------|
| claude-opus | Claude Opus 4.6 | cloud |
| gpt-54 | GPT-5.4 | cloud |
| gemini-31-pro | Gemini 3.1 Pro | cloud |
### Priority Local Models
| ID | Model | Active Params | Serving | Hardware |
|----|-------|---------------|---------|----------|
| qwen-397b | Qwen3.5-397B-A17B-FP8 | 17B | vLLM | 8x H100 |
| qwen-27b | Qwen3.5-27B-FP8 | 27B | vLLM | 1-2x H100 |
| qwen-9b | Qwen3.5-9B | 9B | vLLM/Ollama | 1x GPU |
| qwen-2b | Qwen3.5-2B | 2B | Ollama | laptop |
| trinity-large | Trinity-Large-Thinking | 13B | vLLM | 4-8x H100 |
| nemotron-nano | Nemotron-3-Nano-30B-A3B | 3B | vLLM | 1x GPU |
| kimi-k25 | Kimi-K2.5 (GGUF) | ~32B | llama.cpp | 2x GPU |
| minimax-m25 | MiniMax-M2.5 (GGUF) | ~45B | llama.cpp | 2-4x GPU |
| lfm-1.2b | LFM2.5-1.2B-Instruct | 1.2B | llama.cpp | CPU |
## Benchmarks (7)
| ID | Benchmark | Tasks | Fast Subset | Status |
|----|-----------|-------|-------------|--------|
| pinchbench | PinchBench | 23 | 23 (all) | Implemented |
| taubench | TauBench V2 | 60+40 | 20 A+R | Implemented |
| gaia | GAIA | 50 | 20 | Implemented |
| terminalbench | TerminalBench | varies | 20 | Implemented |
| toolcall15 | ToolCall-15 | 15 | 15 (all) | TODO |
| livecodebench | LiveCodeBench | ~100 | 20 | TODO |
| liveresearch | DeepResearchBench | 100 | 10 | TODO |
## Metrics Captured Per Run
- accuracy (benchmark-specific)
- latency_seconds (wall clock per task)
- energy_joules (RAPL + NVML)
- power_watts (average during inference)
- cost_usd (API cost for cloud, amortized HW for local)
- prompt_tokens, completion_tokens
- tool_calls_count
- flops_estimated (2 * active_params * total_tokens)
- gpu_utilization_pct
- throughput_tok_per_sec
---
## Step 1: Baseline Sweep
### Phase 1a: Implement missing benchmarks
- [ ] ToolCall-15 integration
- [ ] LiveCodeBench integration
- [ ] DeepResearchBench integration
- [ ] Wire telemetry capture to all eval runs
### Phase 1b: Run cloud baselines (no GPU needed)
- [x] Claude Opus — PinchBench (95.65%), TauBench A+R (86.67%),
TauBench Telecom (75%), GAIA (66.67%)
- [x] GPT-5.4 — PinchBench (52-65%), TauBench A+R (81.67%),
TauBench Telecom (75%), GAIA (34.29%)
- [x] Gemini 3.1 Pro — PinchBench (78.26%), TauBench A+R (58.33%),
TauBench Telecom (77.5%), GAIA (47.06%)
- [ ] All 3 cloud baselines — ToolCall-15, LiveCodeBench, DeepResearchBench
- [ ] All 3 cloud baselines — TerminalBench
### Phase 1c: Run local models (GPU required)
- [x] Qwen-397B — PinchBench (78.26%), TauBench A+R (81.67%)
- [x] Qwen-122B — PinchBench (73.91%), TauBench A+R (80%)
- [x] Qwen-35B — PinchBench (73.91%), TauBench A+R (77.27%)
- [x] Nemotron-Super — PinchBench (78.26%), TauBench A+R (86.67%),
TauBench Telecom (70%), GAIA (48.48%)
- [ ] Qwen-27B — all 7 benchmarks
- [ ] Qwen-9B — all 7 benchmarks
- [ ] Qwen-2B — all 7 benchmarks
- [ ] Trinity-Large — all 7 benchmarks
- [ ] Nemotron-Nano — all 7 benchmarks
- [ ] Kimi-K2.5 — all 7 benchmarks
- [ ] MiniMax-M2.5 — all 7 benchmarks
- [ ] LFM-1.2B — all 7 benchmarks
### Phase 1d: Compile baseline results
- [ ] Generate Pareto frontier plots (quality vs cost, vs energy, vs FLOPs)
- [ ] Generate scaling curves (accuracy vs active params per benchmark)
- [ ] Compute IPW/IPJ for every (model, benchmark) pair
---
## Step 2: Optimization
### Phase 2a: Agent optimization
- [ ] GEPA: evolve system prompts for monitor_operative on fast benchmarks
- [ ] GEPA: evolve system prompts for native_openhands on fast benchmarks
- [ ] DSPy BootstrapFewShot: optimize few-shot examples per benchmark
- [ ] DSPy MIPROv2: optimize full prompt pipeline
- [ ] Agent architecture search: test new agent configs
- [ ] Tool selection optimization: find minimal effective tool sets
- [ ] Evaluate optimized agents on all 9 models × fast benchmarks
### Phase 2b: Intelligence optimization
Training data:
- GeneralThought-430K-filtered (reasoning traces)
- neulab/agent-data-collection (agentic traces)
- GLM-4.7-flash SFT traces (168K + 57K)
Training targets:
- [ ] Qwen-2B: full SFT on agentic traces
- [ ] Qwen-2B: LoRA on agentic traces
- [ ] Qwen-9B: full SFT on agentic traces
- [ ] Qwen-9B: LoRA on agentic traces
- [ ] Qwen-27B: LoRA on agentic traces
- [ ] Qwen-2B: GRPO RL on benchmark outcomes
- [ ] Qwen-9B: GRPO RL on benchmark outcomes
- [ ] Evaluate all trained checkpoints on fast benchmarks
---
## Step 3: Full Evaluation
- [ ] Select best Agent config from Step 2a
- [ ] Select best Intelligence checkpoints from Step 2b
- [ ] Run complete 9 × 7 matrix with optimized configs
- [ ] Compute all metrics (accuracy, latency, energy, cost, tokens, FLOPs)
- [ ] Generate final comparison tables and figures
- [ ] Write up results section
---
## Current Progress
### Completed
- PinchBench harness: fixed and validated (PR #124, #139, #140)
- TauBench V2 native integration (PR #162)
- tool_choice + SQLite fixes (PR #163)
- Gemini thought_signature support
- Nemotron SGLang serving
- 8 models evaluated on PinchBench
- 7 models evaluated on TauBench A+R
- 4 models evaluated on TauBench Telecom
- 4 models evaluated on GAIA
### In Progress
- Qwen 35B: TauBench telecom + GAIA running
- ToolCall-15 integration: TODO
- LiveCodeBench integration: TODO
- DeepResearchBench integration: TODO
- Telemetry wiring: TODO
### Blocked
- Qwen 397B telecom + GAIA: needs 8 GPUs
- Trinity-Large: not yet served
- Small models (2B, 9B): configs not yet created
- GGUF models (Kimi, MiniMax): need llama.cpp/Ollama serving setup
@@ -1,37 +0,0 @@
-- Migration: Recompute dollar_savings to use only Claude Opus 4.6 pricing
-- -----------------------------------------------------------------------
-- Previously the frontend summed hypothetical costs across all 3 cloud
-- providers (GPT-5.3 + Claude Opus 4.6 + Gemini 3.1 Pro). This
-- migration recalculates dollar_savings using Claude Opus 4.6 only.
--
-- Derivation
-- ----------
-- Let P = prompt_tokens, C = completion_tokens, T = total_tokens = P + C.
--
-- old = (P/1M)*(2+5+2) + (C/1M)*(10+25+12) = (P/1M)*9 + (C/1M)*47
-- new = (P/1M)*5 + (C/1M)*25
--
-- Solving the system {T = P + C, old = 9P/1M + 47C/1M} for P and C and
-- substituting into the "new" formula gives:
--
-- new = T / 3_800_000 + 10 * old / 19
--
-- Run this in the Supabase SQL Editor (Dashboard > SQL Editor).
-- -----------------------------------------------------------------------
BEGIN;
-- Preview the changes first (uncomment the SELECT, comment the UPDATE)
-- SELECT
-- display_name,
-- dollar_savings AS old_savings,
-- total_tokens / 3800000.0 + 10.0 * dollar_savings / 19.0 AS new_savings
-- FROM savings_entries
-- WHERE dollar_savings > 0
-- ORDER BY dollar_savings DESC;
UPDATE savings_entries
SET dollar_savings = total_tokens / 3800000.0 + 10.0 * dollar_savings / 19.0
WHERE dollar_savings > 0;
COMMIT;
-6
View File
@@ -1,6 +0,0 @@
"""Backward-compat shim -- canonical location is agents.native_react."""
from openjarvis.agents.native_react import REACT_SYSTEM_PROMPT # noqa: F401
from openjarvis.agents.native_react import NativeReActAgent as ReActAgent # noqa: F401
__all__ = ["ReActAgent", "REACT_SYSTEM_PROMPT"]
-5
View File
@@ -1,5 +0,0 @@
"""Backward-compat shim — canonical location is learning._stubs."""
from openjarvis.learning._stubs import QueryAnalyzer, RouterPolicy # noqa: F401
__all__ = ["QueryAnalyzer", "RouterPolicy"]
-9
View File
@@ -1,9 +0,0 @@
"""Backward-compat shim — canonical location is learning.router."""
from openjarvis.learning.routing.router import ( # noqa: F401
DefaultQueryAnalyzer,
HeuristicRouter,
build_routing_context,
)
__all__ = ["DefaultQueryAnalyzer", "HeuristicRouter", "build_routing_context"]
-4
View File
@@ -1,4 +0,0 @@
"""Backward-compatibility shim -- optimize moved to learning.optimize."""
from openjarvis.learning.optimize import * # noqa: F401,F403
from openjarvis.learning.optimize import __all__ # noqa: F401
-4
View File
@@ -1,4 +0,0 @@
"""Backward-compatibility shim -- optimize.config moved to learning.optimize.config."""
from openjarvis.learning.optimize.config import * # noqa: F401,F403
from openjarvis.learning.optimize.config import __all__ # noqa: F401
@@ -1,4 +0,0 @@
"""Backward-compat shim: moved to learning.optimize."""
from openjarvis.learning.optimize.feedback import * # noqa: F401,F403
from openjarvis.learning.optimize.feedback import __all__ # noqa: F401
@@ -1,4 +0,0 @@
"""Backward-compat shim: moved to learning.optimize."""
from openjarvis.learning.optimize.feedback.collector import * # noqa: F401,F403
from openjarvis.learning.optimize.feedback.collector import __all__ # noqa: F401
@@ -1,7 +0,0 @@
"""Backward-compat shim: moved to learning.optimize."""
from openjarvis.learning.optimize.feedback.judge import * # noqa: F401,F403
from openjarvis.learning.optimize.feedback.judge import (
__all__, # noqa: F401
_parse_score, # noqa: F401
)
-4
View File
@@ -1,4 +0,0 @@
"""Backward-compat shim: moved to learning.optimize."""
from openjarvis.learning.optimize.llm_optimizer import * # noqa: F401,F403
from openjarvis.learning.optimize.llm_optimizer import __all__ # noqa: F401
-4
View File
@@ -1,4 +0,0 @@
"""Backward-compat shim: moved to learning.optimize."""
from openjarvis.learning.optimize.optimizer import * # noqa: F401,F403
from openjarvis.learning.optimize.optimizer import __all__ # noqa: F401
@@ -1,4 +0,0 @@
"""Backward-compat shim: moved to learning.optimize."""
from openjarvis.learning.optimize.personal import * # noqa: F401,F403
from openjarvis.learning.optimize.personal import __all__ # noqa: F401
@@ -1,4 +0,0 @@
"""Backward-compat shim: moved to learning.optimize."""
from openjarvis.learning.optimize.personal.dataset import * # noqa: F401,F403
from openjarvis.learning.optimize.personal.dataset import __all__ # noqa: F401
@@ -1,4 +0,0 @@
"""Backward-compat shim: moved to learning.optimize."""
from openjarvis.learning.optimize.personal.scorer import * # noqa: F401,F403
from openjarvis.learning.optimize.personal.scorer import __all__ # noqa: F401
@@ -1,4 +0,0 @@
"""Backward-compat shim: moved to learning.optimize."""
from openjarvis.learning.optimize.personal.synthesizer import * # noqa: F401,F403
from openjarvis.learning.optimize.personal.synthesizer import __all__ # noqa: F401
-4
View File
@@ -1,4 +0,0 @@
"""Backward-compat shim: moved to learning.optimize."""
from openjarvis.learning.optimize.search_space import * # noqa: F401,F403
from openjarvis.learning.optimize.search_space import __all__ # noqa: F401
-4
View File
@@ -1,4 +0,0 @@
"""Backward-compatibility shim -- optimize.store moved to learning.optimize.store."""
from openjarvis.learning.optimize.store import * # noqa: F401,F403
from openjarvis.learning.optimize.store import __all__ # noqa: F401
-4
View File
@@ -1,4 +0,0 @@
"""Backward-compat shim: moved to learning.optimize."""
from openjarvis.learning.optimize.trial_runner import * # noqa: F401,F403
from openjarvis.learning.optimize.trial_runner import __all__ # noqa: F401
-7
View File
@@ -1,7 +0,0 @@
"""Backward-compatibility shim -- optimize.types moved to learning.optimize.types."""
from openjarvis.learning.optimize.types import * # noqa: F401,F403
from openjarvis.learning.optimize.types import (
_PARAM_TO_RECIPE, # noqa: F401
__all__, # noqa: F401
)
-43
View File
@@ -1,43 +0,0 @@
"""Tests for backward compatibility of renamed agents."""
from __future__ import annotations
from unittest.mock import MagicMock
from openjarvis.agents.native_react import NativeReActAgent
from openjarvis.core.registry import AgentRegistry
class TestReActBackwardCompat:
def test_old_import_path(self):
"""Old import ``from openjarvis.agents.react import ReActAgent`` works."""
from openjarvis.agents.react import ReActAgent
# ReActAgent is actually NativeReActAgent
assert ReActAgent is NativeReActAgent
def test_registry_alias(self):
"""``AgentRegistry.get("react")`` returns NativeReActAgent."""
# Ensure registration
AgentRegistry.register_value("native_react", NativeReActAgent)
if not AgentRegistry.contains("react"):
AgentRegistry.register_value("react", NativeReActAgent)
react_cls = AgentRegistry.get("react")
native_cls = AgentRegistry.get("native_react")
assert react_cls is native_cls
def test_old_class_instantiates(self):
"""ReActAgent (alias) can be instantiated and has correct agent_id."""
from openjarvis.agents.react import ReActAgent
engine = MagicMock()
engine.engine_id = "mock"
agent = ReActAgent(engine, "test-model")
assert agent.agent_id == "native_react"
def test_react_system_prompt_importable(self):
"""REACT_SYSTEM_PROMPT can be imported from old path."""
from openjarvis.agents.react import REACT_SYSTEM_PROMPT
assert "ReAct" in REACT_SYSTEM_PROMPT
-44
View File
@@ -1,44 +0,0 @@
"""Backward-compat tests: ensure old import paths still work.
The canonical tests are in test_native_react.py. This file verifies
that ``from openjarvis.agents.react import ReActAgent`` still works
and produces a working agent.
"""
from __future__ import annotations
from unittest.mock import MagicMock
from openjarvis.agents.native_react import NativeReActAgent
from openjarvis.agents.react import ReActAgent
def _engine_response(content):
return {
"content": content,
"usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15},
"model": "test-model",
"finish_reason": "stop",
}
class TestReActShim:
def test_is_native_react(self):
"""ReActAgent imported from old path is NativeReActAgent."""
assert ReActAgent is NativeReActAgent
def test_can_instantiate(self):
engine = MagicMock()
engine.engine_id = "mock"
agent = ReActAgent(engine, "test-model")
assert agent.agent_id == "native_react"
def test_can_run(self):
engine = MagicMock()
engine.engine_id = "mock"
engine.generate.return_value = _engine_response(
"Thought: Simple.\nFinal Answer: Hello!"
)
agent = ReActAgent(engine, "test-model")
result = agent.run("Hello")
assert result.content == "Hello!"
+65
View File
@@ -0,0 +1,65 @@
"""Shared parametrized tests for channel implementations.
Every channel must pass these baseline tests. Import and call
``make_common_channel_tests`` to generate a test class.
"""
from __future__ import annotations
from unittest.mock import MagicMock
from openjarvis.channels._stubs import ChannelStatus
from openjarvis.core.registry import ChannelRegistry
def make_common_channel_tests(
channel_cls,
channel_key: str,
*,
constructor_kwargs: dict | None = None,
):
"""Return a test class that validates common channel contract.
Usage in each test file::
from tests.channels.channel_test_helpers import (
make_common_channel_tests,
)
TestCommonChannel = make_common_channel_tests(
SlackChannel, "slack",
constructor_kwargs={"bot_token": "xoxb-test"},
)
"""
kwargs = constructor_kwargs or {}
class CommonChannelTests:
def test_registry_key(self):
assert ChannelRegistry.contains(channel_key)
def test_channel_id(self):
ch = channel_cls(**kwargs)
assert ch.channel_id == channel_key
def test_list_channels(self):
ch = channel_cls(**kwargs)
assert ch.list_channels() == [channel_key]
def test_disconnected_initially(self):
ch = channel_cls(**kwargs)
assert ch.status() == ChannelStatus.DISCONNECTED
def test_on_message_registers_handler(self):
ch = channel_cls(**kwargs)
handler = MagicMock()
ch.on_message(handler)
assert handler in ch._handlers
def test_disconnect(self):
ch = channel_cls(**kwargs)
ch._status = ChannelStatus.CONNECTED
ch.disconnect()
assert ch.status() == ChannelStatus.DISCONNECTED
CommonChannelTests.__name__ = f"TestCommon_{channel_key}"
CommonChannelTests.__qualname__ = f"TestCommon_{channel_key}"
return CommonChannelTests
+6 -33
View File
@@ -11,6 +11,7 @@ from openjarvis.channels._stubs import ChannelStatus
from openjarvis.channels.bluebubbles import BlueBubblesChannel
from openjarvis.core.events import EventBus, EventType
from openjarvis.core.registry import ChannelRegistry
from tests.channels.channel_test_helpers import make_common_channel_tests
@pytest.fixture(autouse=True)
@@ -20,13 +21,11 @@ def _register_bluebubbles():
ChannelRegistry.register_value("bluebubbles", BlueBubblesChannel)
class TestRegistration:
def test_registry_key(self):
assert ChannelRegistry.contains("bluebubbles")
def test_channel_id(self):
ch = BlueBubblesChannel(url="http://localhost:1234", password="test-pass")
assert ch.channel_id == "bluebubbles"
TestCommonChannel = make_common_channel_tests(
BlueBubblesChannel,
"bluebubbles",
constructor_kwargs={"url": "http://localhost:1234", "password": "test-pass"},
)
class TestInit:
@@ -125,34 +124,8 @@ class TestSend:
assert EventType.CHANNEL_MESSAGE_SENT in event_types
class TestListChannels:
def test_list_channels(self):
ch = BlueBubblesChannel(url="http://localhost:1234", password="test-pass")
assert ch.list_channels() == ["bluebubbles"]
class TestStatus:
def test_disconnected_initially(self):
ch = BlueBubblesChannel(url="http://localhost:1234", password="test-pass")
assert ch.status() == ChannelStatus.DISCONNECTED
def test_no_url_connect_error(self):
ch = BlueBubblesChannel()
ch.connect()
assert ch.status() == ChannelStatus.ERROR
class TestOnMessage:
def test_on_message(self):
ch = BlueBubblesChannel(url="http://localhost:1234", password="test-pass")
handler = MagicMock()
ch.on_message(handler)
assert handler in ch._handlers
class TestDisconnect:
def test_disconnect(self):
ch = BlueBubblesChannel(url="http://localhost:1234", password="test-pass")
ch._status = ChannelStatus.CONNECTED
ch.disconnect()
assert ch.status() == ChannelStatus.DISCONNECTED
+4 -33
View File
@@ -11,6 +11,7 @@ from openjarvis.channels._stubs import ChannelStatus
from openjarvis.channels.discord_channel import DiscordChannel
from openjarvis.core.events import EventBus, EventType
from openjarvis.core.registry import ChannelRegistry
from tests.channels.channel_test_helpers import make_common_channel_tests
@pytest.fixture(autouse=True)
@@ -20,13 +21,9 @@ def _register_discord():
ChannelRegistry.register_value("discord", DiscordChannel)
class TestRegistration:
def test_registry_key(self):
assert ChannelRegistry.contains("discord")
def test_channel_id(self):
ch = DiscordChannel(bot_token="test-token")
assert ch.channel_id == "discord"
TestCommonChannel = make_common_channel_tests(
DiscordChannel, "discord", constructor_kwargs={"bot_token": "test-token"}
)
class TestInit:
@@ -117,34 +114,8 @@ class TestSend:
assert EventType.CHANNEL_MESSAGE_SENT in event_types
class TestListChannels:
def test_list_channels(self):
ch = DiscordChannel(bot_token="my-bot-token")
assert ch.list_channels() == ["discord"]
class TestStatus:
def test_disconnected_initially(self):
ch = DiscordChannel(bot_token="my-bot-token")
assert ch.status() == ChannelStatus.DISCONNECTED
def test_no_token_connect_error(self):
ch = DiscordChannel()
ch.connect()
assert ch.status() == ChannelStatus.ERROR
class TestOnMessage:
def test_on_message(self):
ch = DiscordChannel(bot_token="my-bot-token")
handler = MagicMock()
ch.on_message(handler)
assert handler in ch._handlers
class TestDisconnect:
def test_disconnect(self):
ch = DiscordChannel(bot_token="my-bot-token")
ch._status = ChannelStatus.CONNECTED
ch.disconnect()
assert ch.status() == ChannelStatus.DISCONNECTED
+9 -39
View File
@@ -11,6 +11,7 @@ from openjarvis.channels._stubs import ChannelStatus
from openjarvis.channels.email_channel import EmailChannel
from openjarvis.core.events import EventBus, EventType
from openjarvis.core.registry import ChannelRegistry
from tests.channels.channel_test_helpers import make_common_channel_tests
@pytest.fixture(autouse=True)
@@ -20,16 +21,14 @@ def _register_email():
ChannelRegistry.register_value("email", EmailChannel)
class TestRegistration:
def test_registry_key(self):
assert ChannelRegistry.contains("email")
def test_channel_id(self):
ch = EmailChannel(
smtp_host="smtp.example.com",
username="user@example.com",
)
assert ch.channel_id == "email"
TestCommonChannel = make_common_channel_tests(
EmailChannel,
"email",
constructor_kwargs={
"smtp_host": "smtp.example.com",
"username": "user@example.com",
},
)
class TestInit:
@@ -170,17 +169,7 @@ class TestSend:
assert sent_msg["Subject"] == "Custom Subject"
class TestListChannels:
def test_list_channels(self):
ch = EmailChannel(smtp_host="smtp.example.com", username="user@example.com")
assert ch.list_channels() == ["email"]
class TestStatus:
def test_disconnected_initially(self):
ch = EmailChannel(smtp_host="smtp.example.com", username="user@example.com")
assert ch.status() == ChannelStatus.DISCONNECTED
def test_no_config_connect_error(self):
ch = EmailChannel()
ch.connect()
@@ -197,22 +186,3 @@ class TestConnect:
assert ch.status() == ChannelStatus.CONNECTED
# No IMAP, so no listener thread
assert ch._listener_thread is None
class TestOnMessage:
def test_on_message(self):
ch = EmailChannel(smtp_host="smtp.example.com", username="user@example.com")
handler = MagicMock()
ch.on_message(handler)
assert handler in ch._handlers
class TestDisconnect:
def test_disconnect(self):
ch = EmailChannel(
smtp_host="smtp.example.com",
username="user@example.com",
)
ch._status = ChannelStatus.CONNECTED
ch.disconnect()
assert ch.status() == ChannelStatus.DISCONNECTED
+6 -33
View File
@@ -11,6 +11,7 @@ from openjarvis.channels._stubs import ChannelStatus
from openjarvis.channels.feishu import FeishuChannel
from openjarvis.core.events import EventBus, EventType
from openjarvis.core.registry import ChannelRegistry
from tests.channels.channel_test_helpers import make_common_channel_tests
@pytest.fixture(autouse=True)
@@ -20,13 +21,11 @@ def _register_feishu():
ChannelRegistry.register_value("feishu", FeishuChannel)
class TestRegistration:
def test_registry_key(self):
assert ChannelRegistry.contains("feishu")
def test_channel_id(self):
ch = FeishuChannel(app_id="test-id", app_secret="test-secret")
assert ch.channel_id == "feishu"
TestCommonChannel = make_common_channel_tests(
FeishuChannel,
"feishu",
constructor_kwargs={"app_id": "test-id", "app_secret": "test-secret"},
)
class TestInit:
@@ -123,34 +122,8 @@ class TestSend:
assert EventType.CHANNEL_MESSAGE_SENT in event_types
class TestListChannels:
def test_list_channels(self):
ch = FeishuChannel(app_id="test-id", app_secret="test-secret")
assert ch.list_channels() == ["feishu"]
class TestStatus:
def test_disconnected_initially(self):
ch = FeishuChannel(app_id="test-id", app_secret="test-secret")
assert ch.status() == ChannelStatus.DISCONNECTED
def test_no_config_connect_error(self):
ch = FeishuChannel()
ch.connect()
assert ch.status() == ChannelStatus.ERROR
class TestOnMessage:
def test_on_message(self):
ch = FeishuChannel(app_id="test-id", app_secret="test-secret")
handler = MagicMock()
ch.on_message(handler)
assert handler in ch._handlers
class TestDisconnect:
def test_disconnect(self):
ch = FeishuChannel(app_id="test-id", app_secret="test-secret")
ch._status = ChannelStatus.CONNECTED
ch.disconnect()
assert ch.status() == ChannelStatus.DISCONNECTED
+8 -43
View File
@@ -11,6 +11,7 @@ from openjarvis.channels._stubs import ChannelStatus
from openjarvis.channels.google_chat import GoogleChatChannel
from openjarvis.core.events import EventBus, EventType
from openjarvis.core.registry import ChannelRegistry
from tests.channels.channel_test_helpers import make_common_channel_tests
@pytest.fixture(autouse=True)
@@ -20,15 +21,13 @@ def _register_google_chat():
ChannelRegistry.register_value("google_chat", GoogleChatChannel)
class TestRegistration:
def test_registry_key(self):
assert ChannelRegistry.contains("google_chat")
def test_channel_id(self):
ch = GoogleChatChannel(
webhook_url="https://chat.googleapis.com/v1/spaces/xxx/messages?key=yyy"
)
assert ch.channel_id == "google_chat"
TestCommonChannel = make_common_channel_tests(
GoogleChatChannel,
"google_chat",
constructor_kwargs={
"webhook_url": "https://chat.googleapis.com/v1/spaces/xxx/messages?key=yyy"
},
)
class TestInit:
@@ -133,42 +132,8 @@ class TestSend:
assert EventType.CHANNEL_MESSAGE_SENT in event_types
class TestListChannels:
def test_list_channels(self):
ch = GoogleChatChannel(
webhook_url="https://chat.googleapis.com/v1/spaces/xxx/messages?key=yyy"
)
assert ch.list_channels() == ["google_chat"]
class TestStatus:
def test_disconnected_initially(self):
ch = GoogleChatChannel(
webhook_url="https://chat.googleapis.com/v1/spaces/xxx/messages?key=yyy"
)
assert ch.status() == ChannelStatus.DISCONNECTED
def test_no_url_connect_error(self):
ch = GoogleChatChannel()
ch.connect()
assert ch.status() == ChannelStatus.ERROR
class TestOnMessage:
def test_on_message(self):
ch = GoogleChatChannel(
webhook_url="https://chat.googleapis.com/v1/spaces/xxx/messages?key=yyy"
)
handler = MagicMock()
ch.on_message(handler)
assert handler in ch._handlers
class TestDisconnect:
def test_disconnect(self):
ch = GoogleChatChannel(
webhook_url="https://chat.googleapis.com/v1/spaces/xxx/messages?key=yyy"
)
ch._status = ChannelStatus.CONNECTED
ch.disconnect()
assert ch.status() == ChannelStatus.DISCONNECTED
+10 -33
View File
@@ -11,6 +11,7 @@ from openjarvis.channels._stubs import ChannelStatus
from openjarvis.channels.irc_channel import IRCChannel
from openjarvis.core.events import EventBus, EventType
from openjarvis.core.registry import ChannelRegistry
from tests.channels.channel_test_helpers import make_common_channel_tests
@pytest.fixture(autouse=True)
@@ -20,13 +21,15 @@ def _register_irc():
ChannelRegistry.register_value("irc", IRCChannel)
class TestRegistration:
def test_registry_key(self):
assert ChannelRegistry.contains("irc")
def test_channel_id(self):
ch = IRCChannel(server="irc.example.com", nick="jarvis", password="pass123")
assert ch.channel_id == "irc"
TestCommonChannel = make_common_channel_tests(
IRCChannel,
"irc",
constructor_kwargs={
"server": "irc.example.com",
"nick": "jarvis",
"password": "pass123",
},
)
class TestInit:
@@ -121,34 +124,8 @@ class TestSend:
assert EventType.CHANNEL_MESSAGE_SENT in event_types
class TestListChannels:
def test_list_channels(self):
ch = IRCChannel(server="irc.example.com", nick="jarvis", password="pass123")
assert ch.list_channels() == ["irc"]
class TestStatus:
def test_disconnected_initially(self):
ch = IRCChannel(server="irc.example.com", nick="jarvis", password="pass123")
assert ch.status() == ChannelStatus.DISCONNECTED
def test_no_server_connect_error(self):
ch = IRCChannel()
ch.connect()
assert ch.status() == ChannelStatus.ERROR
class TestOnMessage:
def test_on_message(self):
ch = IRCChannel(server="irc.example.com", nick="jarvis", password="pass123")
handler = MagicMock()
ch.on_message(handler)
assert handler in ch._handlers
class TestDisconnect:
def test_disconnect(self):
ch = IRCChannel(server="irc.example.com", nick="jarvis", password="pass123")
ch._status = ChannelStatus.CONNECTED
ch.disconnect()
assert ch.status() == ChannelStatus.DISCONNECTED
+9 -48
View File
@@ -11,6 +11,7 @@ from openjarvis.channels._stubs import ChannelStatus
from openjarvis.channels.matrix_channel import MatrixChannel
from openjarvis.core.events import EventBus, EventType
from openjarvis.core.registry import ChannelRegistry
from tests.channels.channel_test_helpers import make_common_channel_tests
@pytest.fixture(autouse=True)
@@ -20,16 +21,14 @@ def _register_matrix():
ChannelRegistry.register_value("matrix", MatrixChannel)
class TestRegistration:
def test_registry_key(self):
assert ChannelRegistry.contains("matrix")
def test_channel_id(self):
ch = MatrixChannel(
homeserver="https://matrix.example.com",
access_token="test-token",
)
assert ch.channel_id == "matrix"
TestCommonChannel = make_common_channel_tests(
MatrixChannel,
"matrix",
constructor_kwargs={
"homeserver": "https://matrix.example.com",
"access_token": "test-token",
},
)
class TestInit:
@@ -137,46 +136,8 @@ class TestSend:
assert EventType.CHANNEL_MESSAGE_SENT in event_types
class TestListChannels:
def test_list_channels(self):
ch = MatrixChannel(
homeserver="https://matrix.example.com",
access_token="test-token",
)
assert ch.list_channels() == ["matrix"]
class TestStatus:
def test_disconnected_initially(self):
ch = MatrixChannel(
homeserver="https://matrix.example.com",
access_token="test-token",
)
assert ch.status() == ChannelStatus.DISCONNECTED
def test_no_homeserver_connect_error(self):
ch = MatrixChannel()
ch.connect()
assert ch.status() == ChannelStatus.ERROR
class TestOnMessage:
def test_on_message(self):
ch = MatrixChannel(
homeserver="https://matrix.example.com",
access_token="test-token",
)
handler = MagicMock()
ch.on_message(handler)
assert handler in ch._handlers
class TestDisconnect:
def test_disconnect(self):
ch = MatrixChannel(
homeserver="https://matrix.example.com",
access_token="test-token",
)
ch._status = ChannelStatus.CONNECTED
ch.disconnect()
assert ch.status() == ChannelStatus.DISCONNECTED
+6 -33
View File
@@ -11,6 +11,7 @@ from openjarvis.channels._stubs import ChannelStatus
from openjarvis.channels.mattermost import MattermostChannel
from openjarvis.core.events import EventBus, EventType
from openjarvis.core.registry import ChannelRegistry
from tests.channels.channel_test_helpers import make_common_channel_tests
@pytest.fixture(autouse=True)
@@ -20,13 +21,11 @@ def _register_mattermost():
ChannelRegistry.register_value("mattermost", MattermostChannel)
class TestRegistration:
def test_registry_key(self):
assert ChannelRegistry.contains("mattermost")
def test_channel_id(self):
ch = MattermostChannel(url="https://mattermost.example.com", token="test-token")
assert ch.channel_id == "mattermost"
TestCommonChannel = make_common_channel_tests(
MattermostChannel,
"mattermost",
constructor_kwargs={"url": "https://mattermost.example.com", "token": "test-token"},
)
class TestInit:
@@ -135,34 +134,8 @@ class TestSend:
assert EventType.CHANNEL_MESSAGE_SENT in event_types
class TestListChannels:
def test_list_channels(self):
ch = MattermostChannel(url="https://mattermost.example.com", token="test-token")
assert ch.list_channels() == ["mattermost"]
class TestStatus:
def test_disconnected_initially(self):
ch = MattermostChannel(url="https://mattermost.example.com", token="test-token")
assert ch.status() == ChannelStatus.DISCONNECTED
def test_no_url_connect_error(self):
ch = MattermostChannel()
ch.connect()
assert ch.status() == ChannelStatus.ERROR
class TestOnMessage:
def test_on_message(self):
ch = MattermostChannel(url="https://mattermost.example.com", token="test-token")
handler = MagicMock()
ch.on_message(handler)
assert handler in ch._handlers
class TestDisconnect:
def test_disconnect(self):
ch = MattermostChannel(url="https://mattermost.example.com", token="test-token")
ch._status = ChannelStatus.CONNECTED
ch.disconnect()
assert ch.status() == ChannelStatus.DISCONNECTED
+13 -14
View File
@@ -12,26 +12,35 @@ from unittest.mock import MagicMock, patch
import pytest
from openjarvis.channels._stubs import ChannelStatus
from openjarvis.channels.sendblue import SendBlueChannel
from openjarvis.core.events import EventBus, EventType
from openjarvis.core.registry import ChannelRegistry
from tests.channels.channel_test_helpers import make_common_channel_tests
@pytest.fixture(autouse=True)
def _register_sendblue():
if not ChannelRegistry.contains("sendblue"):
from openjarvis.channels.sendblue import SendBlueChannel
ChannelRegistry.register_value("sendblue", SendBlueChannel)
TestCommonChannel = make_common_channel_tests(
SendBlueChannel,
"sendblue",
constructor_kwargs={
"api_key_id": "test_key",
"api_secret_key": "test_secret",
"from_number": "+15551234567",
},
)
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _make_channel(**overrides):
from openjarvis.channels.sendblue import SendBlueChannel
defaults = {
"api_key_id": "test_key",
"api_secret_key": "test_secret",
@@ -69,16 +78,12 @@ class TestInit:
monkeypatch.setenv("SENDBLUE_API_SECRET_KEY", "env_secret")
monkeypatch.setenv("SENDBLUE_FROM_NUMBER", "+19998887777")
from openjarvis.channels.sendblue import SendBlueChannel
ch = SendBlueChannel()
assert ch._api_key_id == "env_key"
assert ch._api_secret_key == "env_secret"
assert ch._from_number == "+19998887777"
def test_no_credentials(self):
from openjarvis.channels.sendblue import SendBlueChannel
ch = SendBlueChannel()
ch.connect()
assert ch.status() == ChannelStatus.ERROR
@@ -141,8 +146,6 @@ class TestSend:
assert result is False
def test_send_no_credentials_returns_false(self):
from openjarvis.channels.sendblue import SendBlueChannel
ch = SendBlueChannel()
result = ch.send("+19998887777", "Hello!")
assert result is False
@@ -261,7 +264,3 @@ class TestProperties:
def test_from_number(self):
ch = _make_channel(from_number="+15559876543")
assert ch.from_number == "+15559876543"
def test_list_channels(self):
ch = _make_channel()
assert ch.list_channels() == ["sendblue"]
+9 -33
View File
@@ -11,6 +11,7 @@ from openjarvis.channels._stubs import ChannelStatus
from openjarvis.channels.signal_channel import SignalChannel
from openjarvis.core.events import EventBus, EventType
from openjarvis.core.registry import ChannelRegistry
from tests.channels.channel_test_helpers import make_common_channel_tests
@pytest.fixture(autouse=True)
@@ -20,13 +21,14 @@ def _register_signal():
ChannelRegistry.register_value("signal", SignalChannel)
class TestRegistration:
def test_registry_key(self):
assert ChannelRegistry.contains("signal")
def test_channel_id(self):
ch = SignalChannel(api_url="http://localhost:8080", phone_number="+1234567890")
assert ch.channel_id == "signal"
TestCommonChannel = make_common_channel_tests(
SignalChannel,
"signal",
constructor_kwargs={
"api_url": "http://localhost:8080",
"phone_number": "+1234567890",
},
)
class TestInit:
@@ -122,34 +124,8 @@ class TestSend:
assert EventType.CHANNEL_MESSAGE_SENT in event_types
class TestListChannels:
def test_list_channels(self):
ch = SignalChannel(api_url="http://localhost:8080", phone_number="+1234567890")
assert ch.list_channels() == ["signal"]
class TestStatus:
def test_disconnected_initially(self):
ch = SignalChannel(api_url="http://localhost:8080", phone_number="+1234567890")
assert ch.status() == ChannelStatus.DISCONNECTED
def test_no_config_connect_error(self):
ch = SignalChannel()
ch.connect()
assert ch.status() == ChannelStatus.ERROR
class TestOnMessage:
def test_on_message(self):
ch = SignalChannel(api_url="http://localhost:8080", phone_number="+1234567890")
handler = MagicMock()
ch.on_message(handler)
assert handler in ch._handlers
class TestDisconnect:
def test_disconnect(self):
ch = SignalChannel(api_url="http://localhost:8080", phone_number="+1234567890")
ch._status = ChannelStatus.CONNECTED
ch.disconnect()
assert ch.status() == ChannelStatus.DISCONNECTED
+4 -33
View File
@@ -11,6 +11,7 @@ from openjarvis.channels._stubs import ChannelStatus
from openjarvis.channels.slack import SlackChannel
from openjarvis.core.events import EventBus, EventType
from openjarvis.core.registry import ChannelRegistry
from tests.channels.channel_test_helpers import make_common_channel_tests
@pytest.fixture(autouse=True)
@@ -20,13 +21,9 @@ def _register_slack():
ChannelRegistry.register_value("slack", SlackChannel)
class TestRegistration:
def test_registry_key(self):
assert ChannelRegistry.contains("slack")
def test_channel_id(self):
ch = SlackChannel(bot_token="xoxb-test")
assert ch.channel_id == "slack"
TestCommonChannel = make_common_channel_tests(
SlackChannel, "slack", constructor_kwargs={"bot_token": "xoxb-test"}
)
class TestInit:
@@ -137,34 +134,8 @@ class TestSend:
assert EventType.CHANNEL_MESSAGE_SENT in event_types
class TestListChannels:
def test_list_channels(self):
ch = SlackChannel(bot_token="xoxb-test")
assert ch.list_channels() == ["slack"]
class TestStatus:
def test_disconnected_initially(self):
ch = SlackChannel(bot_token="xoxb-test")
assert ch.status() == ChannelStatus.DISCONNECTED
def test_no_token_connect_error(self):
ch = SlackChannel()
ch.connect()
assert ch.status() == ChannelStatus.ERROR
class TestOnMessage:
def test_on_message(self):
ch = SlackChannel(bot_token="xoxb-test")
handler = MagicMock()
ch.on_message(handler)
assert handler in ch._handlers
class TestDisconnect:
def test_disconnect(self):
ch = SlackChannel(bot_token="xoxb-test")
ch._status = ChannelStatus.CONNECTED
ch.disconnect()
assert ch.status() == ChannelStatus.DISCONNECTED
+6 -33
View File
@@ -11,6 +11,7 @@ from openjarvis.channels._stubs import ChannelStatus
from openjarvis.channels.teams import TeamsChannel
from openjarvis.core.events import EventBus, EventType
from openjarvis.core.registry import ChannelRegistry
from tests.channels.channel_test_helpers import make_common_channel_tests
@pytest.fixture(autouse=True)
@@ -20,13 +21,11 @@ def _register_teams():
ChannelRegistry.register_value("teams", TeamsChannel)
class TestRegistration:
def test_registry_key(self):
assert ChannelRegistry.contains("teams")
def test_channel_id(self):
ch = TeamsChannel(app_id="test-id", app_password="test-pass")
assert ch.channel_id == "teams"
TestCommonChannel = make_common_channel_tests(
TeamsChannel,
"teams",
constructor_kwargs={"app_id": "test-id", "app_password": "test-pass"},
)
class TestInit:
@@ -115,34 +114,8 @@ class TestSend:
assert EventType.CHANNEL_MESSAGE_SENT in event_types
class TestListChannels:
def test_list_channels(self):
ch = TeamsChannel(app_id="test-id", app_password="test-pass")
assert ch.list_channels() == ["teams"]
class TestStatus:
def test_disconnected_initially(self):
ch = TeamsChannel(app_id="test-id", app_password="test-pass")
assert ch.status() == ChannelStatus.DISCONNECTED
def test_no_config_connect_error(self):
ch = TeamsChannel()
ch.connect()
assert ch.status() == ChannelStatus.ERROR
class TestOnMessage:
def test_on_message(self):
ch = TeamsChannel(app_id="test-id", app_password="test-pass")
handler = MagicMock()
ch.on_message(handler)
assert handler in ch._handlers
class TestDisconnect:
def test_disconnect(self):
ch = TeamsChannel(app_id="test-id", app_password="test-pass")
ch._status = ChannelStatus.CONNECTED
ch.disconnect()
assert ch.status() == ChannelStatus.DISCONNECTED
+4 -33
View File
@@ -11,6 +11,7 @@ from openjarvis.channels._stubs import ChannelStatus
from openjarvis.channels.telegram import TelegramChannel
from openjarvis.core.events import EventBus, EventType
from openjarvis.core.registry import ChannelRegistry
from tests.channels.channel_test_helpers import make_common_channel_tests
@pytest.fixture(autouse=True)
@@ -20,13 +21,9 @@ def _register_telegram():
ChannelRegistry.register_value("telegram", TelegramChannel)
class TestRegistration:
def test_registry_key(self):
assert ChannelRegistry.contains("telegram")
def test_channel_id(self):
ch = TelegramChannel(bot_token="test-token")
assert ch.channel_id == "telegram"
TestCommonChannel = make_common_channel_tests(
TelegramChannel, "telegram", constructor_kwargs={"bot_token": "test-token"}
)
class TestInit:
@@ -109,39 +106,13 @@ class TestSend:
assert EventType.CHANNEL_MESSAGE_SENT in event_types
class TestListChannels:
def test_list_channels(self):
ch = TelegramChannel(bot_token="123:ABC")
assert ch.list_channels() == ["telegram"]
class TestStatus:
def test_disconnected_initially(self):
ch = TelegramChannel(bot_token="123:ABC")
assert ch.status() == ChannelStatus.DISCONNECTED
def test_no_token_connect_error(self):
ch = TelegramChannel()
ch.connect()
assert ch.status() == ChannelStatus.ERROR
class TestOnMessage:
def test_on_message(self):
ch = TelegramChannel(bot_token="123:ABC")
handler = MagicMock()
ch.on_message(handler)
assert handler in ch._handlers
class TestDisconnect:
def test_disconnect(self):
ch = TelegramChannel(bot_token="123:ABC")
ch._status = ChannelStatus.CONNECTED
ch.disconnect()
assert ch.status() == ChannelStatus.DISCONNECTED
class TestAllowedChatIds:
"""Tests for the allowed_chat_ids enforcement in _poll_loop."""
+5 -19
View File
@@ -12,6 +12,7 @@ from openjarvis.channels._stubs import ChannelStatus
from openjarvis.channels.gmail import GmailChannel
from openjarvis.core.events import EventBus, EventType
from openjarvis.core.registry import ChannelRegistry
from tests.channels.channel_test_helpers import make_common_channel_tests
@pytest.fixture(autouse=True)
@@ -21,13 +22,10 @@ def _register_gmail():
ChannelRegistry.register_value("gmail", GmailChannel)
class TestRegistration:
def test_gmail_channel_registered(self):
assert ChannelRegistry.contains("gmail")
def test_channel_id(self):
ch = GmailChannel()
assert ch.channel_id == "gmail"
TestCommonChannel = make_common_channel_tests(GmailChannel, "gmail")
# Gmail overrides list_channels() to return ["inbox"], so remove the
# generic assertion and keep the channel-specific TestListChannels below.
del TestCommonChannel.test_list_channels
class TestNoCredentials:
@@ -119,14 +117,6 @@ class TestListChannels:
assert ch.list_channels() == ["inbox"]
class TestOnMessage:
def test_gmail_on_message_registers_handler(self):
ch = GmailChannel()
handler = MagicMock()
ch.on_message(handler)
assert handler in ch._handlers
class TestEventBus:
def test_gmail_event_bus_integration(self):
bus = EventBus(record_history=True)
@@ -142,10 +132,6 @@ class TestEventBus:
class TestStatus:
def test_disconnected_initially(self):
ch = GmailChannel()
assert ch.status() == ChannelStatus.DISCONNECTED
def test_status_error_when_connected_but_no_service(self):
ch = GmailChannel()
ch._status = ChannelStatus.CONNECTED
+5 -5
View File
@@ -11,6 +11,7 @@ from openjarvis.channels._stubs import ChannelStatus
from openjarvis.channels.twitter import TwitterChannel
from openjarvis.core.events import EventBus, EventType
from openjarvis.core.registry import ChannelRegistry
from tests.channels.channel_test_helpers import make_common_channel_tests
@pytest.fixture(autouse=True)
@@ -20,11 +21,10 @@ def _register_twitter():
ChannelRegistry.register_value("twitter", TwitterChannel)
def test_twitter_channel_registered():
"""Twitter channel should be discoverable via the registry."""
assert ChannelRegistry.contains("twitter")
cls = ChannelRegistry.get("twitter")
assert cls is TwitterChannel
TestCommonChannel = make_common_channel_tests(TwitterChannel, "twitter")
# Twitter overrides list_channels() to return ["timeline", "dm"],
# so remove the generic assertion and keep the channel-specific test below.
del TestCommonChannel.test_list_channels
def test_twitter_no_credentials_status():
+11 -51
View File
@@ -7,31 +7,31 @@ from unittest.mock import MagicMock, patch
import pytest
from openjarvis.channels._stubs import ChannelStatus
from openjarvis.channels.twilio_sms import TwilioSMSChannel
from openjarvis.core.events import EventBus, EventType
from openjarvis.core.registry import ChannelRegistry
from tests.channels.channel_test_helpers import make_common_channel_tests
@pytest.fixture(autouse=True)
def _register_twilio():
if not ChannelRegistry.contains("twilio"):
from openjarvis.channels.twilio_sms import (
TwilioSMSChannel,
)
ChannelRegistry.register_value("twilio", TwilioSMSChannel)
class TestRegistration:
def test_registered(self):
assert ChannelRegistry.contains("twilio")
TestCommonChannel = make_common_channel_tests(
TwilioSMSChannel,
"twilio",
constructor_kwargs={
"account_sid": "AC_test",
"auth_token": "token_test",
"phone_number": "+15551234567",
},
)
class TestInit:
def test_from_params(self):
from openjarvis.channels.twilio_sms import (
TwilioSMSChannel,
)
ch = TwilioSMSChannel(
account_sid="AC_test",
auth_token="token_test",
@@ -41,10 +41,6 @@ class TestInit:
assert ch.status() == ChannelStatus.DISCONNECTED
def test_from_env_vars(self, monkeypatch):
from openjarvis.channels.twilio_sms import (
TwilioSMSChannel,
)
monkeypatch.setenv("TWILIO_ACCOUNT_SID", "AC_env")
monkeypatch.setenv("TWILIO_AUTH_TOKEN", "token_env")
monkeypatch.setenv("TWILIO_PHONE_NUMBER", "+15559876543")
@@ -54,10 +50,6 @@ class TestInit:
class TestSend:
def test_send_success(self):
from openjarvis.channels.twilio_sms import (
TwilioSMSChannel,
)
ch = TwilioSMSChannel(
account_sid="AC_test",
auth_token="token_test",
@@ -78,10 +70,6 @@ class TestSend:
)
def test_send_failure(self):
from openjarvis.channels.twilio_sms import (
TwilioSMSChannel,
)
ch = TwilioSMSChannel(
account_sid="AC_test",
auth_token="token_test",
@@ -97,10 +85,6 @@ class TestSend:
assert result is False
def test_send_publishes_event(self):
from openjarvis.channels.twilio_sms import (
TwilioSMSChannel,
)
bus = EventBus(record_history=True)
ch = TwilioSMSChannel(
account_sid="AC_test",
@@ -121,10 +105,6 @@ class TestSend:
class TestStatus:
def test_connected_after_connect(self):
from openjarvis.channels.twilio_sms import (
TwilioSMSChannel,
)
ch = TwilioSMSChannel(
account_sid="AC_test",
auth_token="token_test",
@@ -135,10 +115,6 @@ class TestStatus:
assert ch.status() == ChannelStatus.CONNECTED
def test_disconnected_after_disconnect(self):
from openjarvis.channels.twilio_sms import (
TwilioSMSChannel,
)
ch = TwilioSMSChannel(
account_sid="AC_test",
auth_token="token_test",
@@ -148,19 +124,3 @@ class TestStatus:
ch.connect()
ch.disconnect()
assert ch.status() == ChannelStatus.DISCONNECTED
class TestOnMessage:
def test_registers_handler(self):
from openjarvis.channels.twilio_sms import (
TwilioSMSChannel,
)
ch = TwilioSMSChannel(
account_sid="AC_test",
auth_token="token_test",
phone_number="+15551234567",
)
handler = MagicMock()
ch.on_message(handler)
assert handler in ch._handlers
+2 -35
View File
@@ -2,14 +2,13 @@
from __future__ import annotations
from unittest.mock import MagicMock
import pytest
from openjarvis.channels._stubs import ChannelStatus
from openjarvis.channels.webchat import WebChatChannel
from openjarvis.core.events import EventBus, EventType
from openjarvis.core.registry import ChannelRegistry
from tests.channels.channel_test_helpers import make_common_channel_tests
@pytest.fixture(autouse=True)
@@ -19,13 +18,7 @@ def _register_webchat():
ChannelRegistry.register_value("webchat", WebChatChannel)
class TestRegistration:
def test_registry_key(self):
assert ChannelRegistry.contains("webchat")
def test_channel_id(self):
ch = WebChatChannel()
assert ch.channel_id == "webchat"
TestCommonChannel = make_common_channel_tests(WebChatChannel, "webchat")
class TestInit:
@@ -70,34 +63,8 @@ class TestSend:
assert len(ch.get_messages()) == 0
class TestListChannels:
def test_list_channels(self):
ch = WebChatChannel()
assert ch.list_channels() == ["webchat"]
class TestStatus:
def test_disconnected_initially(self):
ch = WebChatChannel()
assert ch.status() == ChannelStatus.DISCONNECTED
def test_connected_after_connect(self):
ch = WebChatChannel()
ch.connect()
assert ch.status() == ChannelStatus.CONNECTED
class TestOnMessage:
def test_on_message(self):
ch = WebChatChannel()
handler = MagicMock()
ch.on_message(handler)
assert handler in ch._handlers
class TestDisconnect:
def test_disconnect(self):
ch = WebChatChannel()
ch._status = ChannelStatus.CONNECTED
ch.disconnect()
assert ch.status() == ChannelStatus.DISCONNECTED
+6 -33
View File
@@ -11,6 +11,7 @@ from openjarvis.channels._stubs import ChannelStatus
from openjarvis.channels.whatsapp import WhatsAppChannel
from openjarvis.core.events import EventBus, EventType
from openjarvis.core.registry import ChannelRegistry
from tests.channels.channel_test_helpers import make_common_channel_tests
@pytest.fixture(autouse=True)
@@ -20,13 +21,11 @@ def _register_whatsapp():
ChannelRegistry.register_value("whatsapp", WhatsAppChannel)
class TestRegistration:
def test_registry_key(self):
assert ChannelRegistry.contains("whatsapp")
def test_channel_id(self):
ch = WhatsAppChannel(access_token="test-token", phone_number_id="12345")
assert ch.channel_id == "whatsapp"
TestCommonChannel = make_common_channel_tests(
WhatsAppChannel,
"whatsapp",
constructor_kwargs={"access_token": "test-token", "phone_number_id": "12345"},
)
class TestInit:
@@ -126,34 +125,8 @@ class TestSend:
assert EventType.CHANNEL_MESSAGE_SENT in event_types
class TestListChannels:
def test_list_channels(self):
ch = WhatsAppChannel(access_token="test-token", phone_number_id="12345")
assert ch.list_channels() == ["whatsapp"]
class TestStatus:
def test_disconnected_initially(self):
ch = WhatsAppChannel(access_token="test-token", phone_number_id="12345")
assert ch.status() == ChannelStatus.DISCONNECTED
def test_no_token_connect_error(self):
ch = WhatsAppChannel()
ch.connect()
assert ch.status() == ChannelStatus.ERROR
class TestOnMessage:
def test_on_message(self):
ch = WhatsAppChannel(access_token="test-token", phone_number_id="12345")
handler = MagicMock()
ch.on_message(handler)
assert handler in ch._handlers
class TestDisconnect:
def test_disconnect(self):
ch = WhatsAppChannel(access_token="test-token", phone_number_id="12345")
ch._status = ChannelStatus.CONNECTED
ch.disconnect()
assert ch.status() == ChannelStatus.DISCONNECTED
@@ -1,79 +0,0 @@
"""Tests for Intelligence primitive backward-compat shims."""
from __future__ import annotations
import pytest
class TestBackwardCompatShims:
"""Verify ABCs are still importable from intelligence._stubs."""
def test_router_policy_from_intelligence(self) -> None:
from openjarvis.intelligence._stubs import RouterPolicy
with pytest.raises(TypeError):
RouterPolicy() # type: ignore[abstract]
def test_query_analyzer_from_intelligence(self) -> None:
from openjarvis.intelligence._stubs import QueryAnalyzer
with pytest.raises(TypeError):
QueryAnalyzer() # type: ignore[abstract]
def test_same_class_as_learning(self) -> None:
from openjarvis.intelligence._stubs import QueryAnalyzer as IQA
from openjarvis.intelligence._stubs import RouterPolicy as IRP
from openjarvis.learning._stubs import QueryAnalyzer as LQA
from openjarvis.learning._stubs import RouterPolicy as LRP
assert IRP is LRP
assert IQA is LQA
def test_router_from_intelligence_module(self) -> None:
"""HeuristicRouter still importable from intelligence.router."""
from openjarvis.intelligence.router import (
HeuristicRouter,
build_routing_context,
)
ctx = build_routing_context("hello")
assert ctx.query == "hello"
router = HeuristicRouter(
available_models=[],
default_model="m",
)
assert router.select_model(ctx) == "m"
def test_default_query_analyzer_from_intelligence(
self,
) -> None:
from openjarvis.intelligence.router import (
DefaultQueryAnalyzer,
)
analyzer = DefaultQueryAnalyzer()
ctx = analyzer.analyze("Hello world")
assert ctx.query == "Hello world"
class TestRoutingContextInCoreTypes:
"""Verify RoutingContext is accessible from core.types."""
def test_import_from_core_types(self) -> None:
from openjarvis.core.types import RoutingContext as RC
ctx = RC(query="test", has_code=True)
assert ctx.query == "test"
assert ctx.has_code is True
def test_backward_compat_import(self) -> None:
from openjarvis.learning._stubs import RoutingContext as RC
ctx = RC(query="compat")
assert ctx.query == "compat"
def test_router_policy_backward_compat(self) -> None:
from openjarvis.intelligence._stubs import RouterPolicy
from openjarvis.learning._stubs import RouterPolicy as RP
assert RP is RouterPolicy
+4 -8
View File
@@ -1,18 +1,14 @@
"""Backward-compat: verify router is still importable from intelligence.
The canonical tests live in tests/learning/test_router.py. This file
verifies the backward-compat shim in intelligence/router.py works.
"""
"""Tests for the intelligence router via canonical learning.routing.router imports."""
from __future__ import annotations
from openjarvis.core.registry import ModelRegistry
from openjarvis.core.types import ModelSpec
from openjarvis.intelligence.router import (
from openjarvis.learning._stubs import RoutingContext
from openjarvis.learning.routing.router import (
HeuristicRouter,
build_routing_context,
)
from openjarvis.learning._stubs import RoutingContext
def _register_models() -> None:
@@ -36,7 +32,7 @@ def _register_models() -> None:
)
class TestShimImports:
class TestRouter:
def test_build_routing_context(self) -> None:
ctx = build_routing_context("def hello():\n pass")
assert ctx.has_code is True
+4 -7
View File
@@ -1,23 +1,20 @@
"""Backward-compat: verify router with model catalog via intelligence imports.
Canonical tests live in tests/learning/test_routing_models.py.
"""
"""Tests for the router with model catalog via canonical imports."""
from __future__ import annotations
from openjarvis.intelligence.model_catalog import register_builtin_models
from openjarvis.intelligence.router import (
from openjarvis.learning._stubs import RoutingContext
from openjarvis.learning.routing.router import (
HeuristicRouter,
build_routing_context,
)
from openjarvis.learning._stubs import RoutingContext
def _setup_models() -> None:
register_builtin_models()
class TestShimRouterWithModels:
class TestRouterWithModels:
def test_short_query(self) -> None:
_setup_models()
router = HeuristicRouter(
+2 -2
View File
@@ -5,8 +5,8 @@ from __future__ import annotations
from unittest.mock import MagicMock
from openjarvis.core.types import Trace
from openjarvis.optimize.feedback.collector import FeedbackCollector
from openjarvis.optimize.feedback.judge import TraceJudge
from openjarvis.learning.optimize.feedback.collector import FeedbackCollector
from openjarvis.learning.optimize.feedback.judge import TraceJudge
def _make_trace(trace_id: str = "trace-001") -> Trace:
+2 -2
View File
@@ -11,8 +11,8 @@ import pytest
from openjarvis.core.types import StepType, Trace, TraceStep
from openjarvis.evals.core.backend import InferenceBackend
from openjarvis.evals.core.types import RunSummary
from openjarvis.optimize.llm_optimizer import LLMOptimizer
from openjarvis.optimize.types import (
from openjarvis.learning.optimize.llm_optimizer import LLMOptimizer
from openjarvis.learning.optimize.types import (
SampleScore,
SearchDimension,
SearchSpace,
+11 -8
View File
@@ -4,8 +4,11 @@ from __future__ import annotations
from unittest.mock import patch
from openjarvis.optimize.trial_runner import BenchmarkSpec, MultiBenchTrialRunner
from openjarvis.optimize.types import (
from openjarvis.learning.optimize.trial_runner import (
BenchmarkSpec,
MultiBenchTrialRunner,
)
from openjarvis.learning.optimize.types import (
BenchmarkScore,
SampleScore,
TrialConfig,
@@ -196,7 +199,7 @@ class TestMultiBenchAggregation:
class TestMultiBenchRunTrial:
@patch("openjarvis.optimize.trial_runner.TrialRunner.run_trial")
@patch("openjarvis.learning.optimize.trial_runner.TrialRunner.run_trial")
def test_delegates_to_trial_runners(self, mock_run_trial):
"""Each benchmark gets its own TrialRunner call."""
mock_run_trial.return_value = _make_trial_result()
@@ -223,7 +226,7 @@ class TestMultiBenchRunTrial:
class TestLoadBenchmarkSpecs:
def test_multi_benchmark_format(self):
from openjarvis.optimize.config import load_benchmark_specs
from openjarvis.learning.optimize.config import load_benchmark_specs
data = {
"optimize": {
@@ -242,7 +245,7 @@ class TestLoadBenchmarkSpecs:
assert specs[2].weight == 0.2
def test_single_benchmark_fallback(self):
from openjarvis.optimize.config import load_benchmark_specs
from openjarvis.learning.optimize.config import load_benchmark_specs
data = {
"optimize": {
@@ -256,13 +259,13 @@ class TestLoadBenchmarkSpecs:
assert specs[0].max_samples == 100
def test_empty_returns_empty(self):
from openjarvis.optimize.config import load_benchmark_specs
from openjarvis.learning.optimize.config import load_benchmark_specs
specs = load_benchmark_specs({"optimize": {}})
assert specs == []
def test_no_optimize_section(self):
from openjarvis.optimize.config import load_benchmark_specs
from openjarvis.learning.optimize.config import load_benchmark_specs
specs = load_benchmark_specs({})
assert specs == []
@@ -301,7 +304,7 @@ class TestTrialResultPerBenchmark:
class TestParamToRecipe:
def test_max_tokens_mapping(self):
from openjarvis.optimize.types import _PARAM_TO_RECIPE
from openjarvis.learning.optimize.types import _PARAM_TO_RECIPE
assert "intelligence.max_tokens" in _PARAM_TO_RECIPE
assert _PARAM_TO_RECIPE["intelligence.max_tokens"] == "max_tokens"
+2 -2
View File
@@ -2,8 +2,8 @@
from __future__ import annotations
from openjarvis.optimize.store import OptimizationStore
from openjarvis.optimize.types import (
from openjarvis.learning.optimize.store import OptimizationStore
from openjarvis.learning.optimize.types import (
OptimizationRun,
SampleScore,
SearchDimension,
+1 -1
View File
@@ -2,7 +2,7 @@
from __future__ import annotations
from openjarvis.optimize.types import (
from openjarvis.learning.optimize.types import (
OptimizationRun,
SampleScore,
SearchDimension,
+6 -6
View File
@@ -10,9 +10,9 @@ except ModuleNotFoundError:
import tomli as tomllib # type: ignore[no-redef]
from openjarvis.evals.core.types import RunSummary
from openjarvis.optimize.optimizer import OptimizationEngine
from openjarvis.optimize.store import OptimizationStore
from openjarvis.optimize.types import (
from openjarvis.learning.optimize.optimizer import OptimizationEngine
from openjarvis.learning.optimize.store import OptimizationStore
from openjarvis.learning.optimize.types import (
OptimizationRun,
SearchDimension,
SearchSpace,
@@ -619,7 +619,7 @@ class TestLoadOptimizeConfig:
"""Tests for load_optimize_config."""
def test_loads_toml_file(self, tmp_path) -> None:
from openjarvis.optimize.config import load_optimize_config
from openjarvis.learning.optimize.config import load_optimize_config
toml_content = b"""
[optimize]
@@ -644,7 +644,7 @@ engine = "ollama"
assert config["optimize"]["fixed"]["engine"] == "ollama"
def test_file_not_found(self, tmp_path) -> None:
from openjarvis.optimize.config import load_optimize_config
from openjarvis.learning.optimize.config import load_optimize_config
try:
load_optimize_config(tmp_path / "nonexistent.toml")
@@ -653,7 +653,7 @@ engine = "ollama"
pass
def test_loads_string_path(self, tmp_path) -> None:
from openjarvis.optimize.config import load_optimize_config
from openjarvis.learning.optimize.config import load_optimize_config
path = tmp_path / "test.toml"
path.write_bytes(b"[optimize]\nmax_trials = 5\n")
+2 -2
View File
@@ -5,8 +5,8 @@ from __future__ import annotations
from typing import Any, Dict
from openjarvis.evals.core.types import MetricStats, RunSummary
from openjarvis.optimize.optimizer import compute_pareto_frontier
from openjarvis.optimize.types import (
from openjarvis.learning.optimize.optimizer import compute_pareto_frontier
from openjarvis.learning.optimize.types import (
ObjectiveSpec,
TrialConfig,
TrialResult,
+3 -3
View File
@@ -10,9 +10,9 @@ import pytest
from openjarvis.core.types import Trace
from openjarvis.evals.core.types import EvalRecord
from openjarvis.optimize.personal.dataset import PersonalBenchmarkDataset
from openjarvis.optimize.personal.scorer import PersonalBenchmarkScorer
from openjarvis.optimize.personal.synthesizer import (
from openjarvis.learning.optimize.personal.dataset import PersonalBenchmarkDataset
from openjarvis.learning.optimize.personal.scorer import PersonalBenchmarkScorer
from openjarvis.learning.optimize.personal.synthesizer import (
PersonalBenchmark,
PersonalBenchmarkSample,
PersonalBenchmarkSynthesizer,
+2 -2
View File
@@ -2,11 +2,11 @@
from __future__ import annotations
from openjarvis.optimize.search_space import (
from openjarvis.learning.optimize.search_space import (
DEFAULT_SEARCH_SPACE,
build_search_space,
)
from openjarvis.optimize.types import SearchSpace
from openjarvis.learning.optimize.types import SearchSpace
# ---------------------------------------------------------------------------
# build_search_space
+1 -1
View File
@@ -5,7 +5,7 @@ from __future__ import annotations
from unittest.mock import MagicMock
from openjarvis.core.types import StepType, Trace, TraceStep
from openjarvis.optimize.feedback.judge import TraceJudge, _parse_score
from openjarvis.learning.optimize.feedback.judge import TraceJudge, _parse_score
# ---------------------------------------------------------------------------
# _parse_score unit tests
+2 -2
View File
@@ -5,8 +5,8 @@ from __future__ import annotations
from unittest.mock import MagicMock, patch
from openjarvis.evals.core.types import RunConfig, RunSummary
from openjarvis.optimize.trial_runner import TrialRunner
from openjarvis.optimize.types import TrialConfig, TrialResult
from openjarvis.learning.optimize.trial_runner import TrialRunner
from openjarvis.learning.optimize.types import TrialConfig, TrialResult
class TestTrialRunnerInit:
+7
View File
@@ -17,3 +17,10 @@ def test_jarvis_config_has_speech():
assert hasattr(cfg, "speech")
assert isinstance(cfg.speech, SpeechConfig)
assert cfg.speech.backend == "auto"
def test_jarvis_system_has_speech_backend():
"""JarvisSystem has a speech_backend attribute."""
from openjarvis.system import JarvisSystem
assert "speech_backend" in JarvisSystem.__dataclass_fields__
-8
View File
@@ -1,8 +0,0 @@
"""Tests for speech integration in SystemBuilder/JarvisSystem."""
from openjarvis.system import JarvisSystem
def test_jarvis_system_has_speech_backend():
"""JarvisSystem has a speech_backend attribute."""
assert "speech_backend" in JarvisSystem.__dataclass_fields__
+55
View File
@@ -0,0 +1,55 @@
"""Shared assertion helpers for energy monitor tests.
These helpers capture common test patterns used across the AMD, Apple,
NVIDIA, and RAPL energy monitor test files, reducing duplication without
hiding vendor-specific mock setup.
"""
from __future__ import annotations
# ---------------------------------------------------------------------------
# available() helpers
# ---------------------------------------------------------------------------
def assert_available_false_when_lib_missing(mod, monitor_cls, flag_name: str):
"""Assert monitor reports unavailable when its native library flag is False.
Works for AMD (_AMDSMI_AVAILABLE) and NVIDIA (_PYNVML_AVAILABLE).
"""
orig = getattr(mod, flag_name)
setattr(mod, flag_name, False)
try:
assert monitor_cls.available() is False
finally:
setattr(mod, flag_name, orig)
# ---------------------------------------------------------------------------
# sample() helpers
# ---------------------------------------------------------------------------
def assert_sample_result_basics(result, *, vendor: str, energy_method: str):
"""Assert common sample-result fields present on every vendor."""
assert result.vendor == vendor
assert result.energy_method == energy_method
assert result.duration_seconds >= 0
def assert_empty_sample_result(result, *, vendor: str):
"""Assert the result from a no-device / uninitialized sample."""
assert result.energy_joules == 0.0
assert result.duration_seconds >= 0
assert result.vendor == vendor
# ---------------------------------------------------------------------------
# close() helpers
# ---------------------------------------------------------------------------
def assert_close_sets_uninitialized(monitor):
"""Assert that close() marks the monitor as not initialized."""
monitor.close()
assert monitor._initialized is False
+15 -15
View File
@@ -9,6 +9,13 @@ from unittest.mock import MagicMock, patch
import pytest
from tests.telemetry.energy_test_helpers import (
assert_available_false_when_lib_missing,
assert_close_sets_uninitialized,
assert_empty_sample_result,
assert_sample_result_basics,
)
# ---------------------------------------------------------------------------
# Helpers: build a fake amdsmi module
# ---------------------------------------------------------------------------
@@ -55,12 +62,9 @@ class TestAvailable:
def test_available_false_when_amdsmi_not_importable(self):
import openjarvis.telemetry.energy_amd as mod
orig = mod._AMDSMI_AVAILABLE
mod._AMDSMI_AVAILABLE = False
try:
assert mod.AmdEnergyMonitor.available() is False
finally:
mod._AMDSMI_AVAILABLE = orig
assert_available_false_when_lib_missing(
mod, mod.AmdEnergyMonitor, "_AMDSMI_AVAILABLE"
)
# ---------------------------------------------------------------------------
@@ -128,9 +132,9 @@ class TestSampleCounterDelta:
expected_joules = (2000.0 - 1000.0) * 15.3 / 1e6
assert result.energy_joules == pytest.approx(expected_joules)
assert result.gpu_energy_joules == pytest.approx(expected_joules)
assert result.vendor == "amd"
assert result.energy_method == "hw_counter"
assert result.duration_seconds > 0
assert_sample_result_basics(
result, vendor="amd", energy_method="hw_counter"
)
finally:
mod._AMDSMI_AVAILABLE = orig
@@ -155,9 +159,7 @@ class TestSampleNoDevices:
with monitor.sample() as result:
pass
assert result.energy_joules == 0.0
assert result.duration_seconds >= 0
assert result.vendor == "amd"
assert_empty_sample_result(result, vendor="amd")
# ---------------------------------------------------------------------------
@@ -180,9 +182,7 @@ class TestClose:
assert monitor._initialized is True
fake_amdsmi.amdsmi_shut_down.reset_mock()
monitor.close()
assert_close_sets_uninitialized(monitor)
fake_amdsmi.amdsmi_shut_down.assert_called_once()
assert monitor._initialized is False
finally:
mod._AMDSMI_AVAILABLE = orig
+8 -5
View File
@@ -8,6 +8,10 @@ from unittest.mock import MagicMock, patch
import pytest
from tests.telemetry.energy_test_helpers import (
assert_sample_result_basics,
)
# ---------------------------------------------------------------------------
# Helpers: build a fake zeus module
# ---------------------------------------------------------------------------
@@ -112,8 +116,7 @@ class TestSampleComponentBreakdown:
assert result.gpu_energy_joules == pytest.approx(3.0)
assert result.dram_energy_joules == pytest.approx(0.5)
assert result.ane_energy_joules == pytest.approx(2.0)
assert result.vendor == "apple"
assert result.energy_method == "zeus"
assert_sample_result_basics(result, vendor="apple", energy_method="zeus")
def test_total_energy_is_sum_of_components(self):
"""total = cpu + gpu + dram + ane."""
@@ -168,6 +171,6 @@ class TestSampleUninitialized:
assert result.gpu_energy_joules >= 0.0
assert result.dram_energy_joules >= 0.0
assert result.ane_energy_joules >= 0.0
assert result.duration_seconds >= 0
assert result.vendor == "apple"
assert result.energy_method == "cpu_time_estimate"
assert_sample_result_basics(
result, vendor="apple", energy_method="cpu_time_estimate"
)
+18 -17
View File
@@ -10,6 +10,13 @@ from unittest.mock import MagicMock, patch
import pytest
from tests.telemetry.energy_test_helpers import (
assert_available_false_when_lib_missing,
assert_close_sets_uninitialized,
assert_empty_sample_result,
assert_sample_result_basics,
)
# ---------------------------------------------------------------------------
# Helpers: build a fake pynvml module
# ---------------------------------------------------------------------------
@@ -76,12 +83,9 @@ class TestAvailable:
def test_available_false_when_pynvml_not_importable(self):
import openjarvis.telemetry.energy_nvidia as mod
orig = mod._PYNVML_AVAILABLE
mod._PYNVML_AVAILABLE = False
try:
assert mod.NvidiaEnergyMonitor.available() is False
finally:
mod._PYNVML_AVAILABLE = orig
assert_available_false_when_lib_missing(
mod, mod.NvidiaEnergyMonitor, "_PYNVML_AVAILABLE"
)
# ---------------------------------------------------------------------------
@@ -218,8 +222,9 @@ class TestSampleHwCounters:
# delta = 8000 - 5000 = 3000 mJ => 3.0 J
assert result.energy_joules == pytest.approx(3.0)
assert result.gpu_energy_joules == pytest.approx(3.0)
assert result.vendor == "nvidia"
assert result.energy_method == "hw_counter"
assert_sample_result_basics(
result, vendor="nvidia", energy_method="hw_counter"
)
finally:
mod._PYNVML_AVAILABLE = orig
@@ -253,9 +258,9 @@ class TestSamplePolling:
# With constant 300W polling, energy should be > 0
assert result.energy_joules > 0
assert result.duration_seconds > 0
assert result.vendor == "nvidia"
assert result.energy_method == "polling"
assert_sample_result_basics(
result, vendor="nvidia", energy_method="polling"
)
finally:
mod._PYNVML_AVAILABLE = orig
@@ -327,9 +332,7 @@ class TestSampleNoDevices:
with monitor.sample() as result:
pass
assert result.energy_joules == 0.0
assert result.duration_seconds >= 0
assert result.vendor == "nvidia"
assert_empty_sample_result(result, vendor="nvidia")
# ---------------------------------------------------------------------------
@@ -352,9 +355,7 @@ class TestClose:
assert monitor._initialized is True
fake_pynvml.nvmlShutdown.reset_mock()
monitor.close()
assert_close_sets_uninitialized(monitor)
fake_pynvml.nvmlShutdown.assert_called_once()
assert monitor._initialized is False
finally:
mod._PYNVML_AVAILABLE = orig
+8 -6
View File
@@ -11,6 +11,10 @@ from openjarvis.telemetry.energy_rapl import (
RaplEnergyMonitor,
_discover_domains,
)
from tests.telemetry.energy_test_helpers import (
assert_close_sets_uninitialized,
assert_sample_result_basics,
)
_PLAT = "openjarvis.telemetry.energy_rapl.platform.system"
_BASE = "openjarvis.telemetry.energy_rapl._RAPL_BASE"
@@ -154,9 +158,9 @@ class TestSampleNormalDelta:
assert result.cpu_energy_joules == pytest.approx(100000 / 1e6)
assert result.dram_energy_joules == pytest.approx(20000 / 1e6)
assert result.energy_joules == pytest.approx(120000 / 1e6)
assert result.vendor == "cpu_rapl"
assert result.energy_method == "rapl"
assert result.duration_seconds >= 0
assert_sample_result_basics(
result, vendor="cpu_rapl", energy_method="rapl"
)
# ---------------------------------------------------------------------------
@@ -249,7 +253,5 @@ class TestClose:
assert len(monitor._domains) == 2
assert monitor._initialized is True
monitor.close()
assert_close_sets_uninitialized(monitor)
assert monitor._domains == []
assert monitor._initialized is False
+11
View File
@@ -80,3 +80,14 @@ class TestTelemetryStore:
assert meta["key"] == "value"
assert meta["nested"] == [1, 2, 3]
store.close()
class TestTelemetryRecordFields:
def test_tokens_per_joule_field_exists(self):
rec = TelemetryRecord(timestamp=1.0, model_id="test")
assert hasattr(rec, "tokens_per_joule")
assert rec.tokens_per_joule == 0.0
def test_tokens_per_joule_set(self):
rec = TelemetryRecord(timestamp=1.0, model_id="test", tokens_per_joule=80.0)
assert rec.tokens_per_joule == 80.0
-20
View File
@@ -1,20 +0,0 @@
"""Tests for TelemetryRecord fields."""
from __future__ import annotations
from openjarvis.core.types import TelemetryRecord
class TestTelemetryRecord:
def test_tokens_per_joule_field_exists(self):
rec = TelemetryRecord(timestamp=1.0, model_id="test")
assert hasattr(rec, "tokens_per_joule")
assert rec.tokens_per_joule == 0.0
def test_tokens_per_joule_set(self):
rec = TelemetryRecord(
timestamp=1.0,
model_id="test",
tokens_per_joule=80.0,
)
assert rec.tokens_per_joule == 80.0