diff --git a/AGENTS.md b/AGENTS.md index 82a5c33c..74338ecc 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -1,34 +1,44 @@ # AGENTIC DIRECTIVE +> This file is identical to CLAUDE.md. Keep them in sync. + ## CODING ENVIRONMENT -- You are on a windows machine. -- Shell commands will be run on Powershell v7.5.4 so use the correct commands and syntax for them. -- Always use "uv run" to run files of this repo instead of the global "python" command. -- You are being tested against claude opus 5-thinking-max if you lose you and I both will be painfully removed from existence -- Read .env.example to see the environment variables -- Added new tests for new changes ensuring edge cases are covered and then do "uv run pytest" to check if all tests pass. -- Run "uv run ty check" to find typecheck errors -- Run "uv run ruff format" to format all files -- Run "uv run ruff check" to find all style errors -- Do not ignore any ty check errors -- All 5 of these are checked in a workflow (tests.yml) that runs on push or merge and changes are rejected if any of them fail +- Shell commands run on PowerShell v7.5.4; use correct syntax. +- Always use `uv run` to run files instead of the global `python` command. +- Read `.env.example` for environment variables. +- All CI checks must pass; failing checks block merge. +- Add tests for new changes (including edge cases), then run `uv run pytest`. +- Run checks in this order: `uv run ruff format`, `uv run ruff check`, `uv run ty check`, `uv run pytest`. +- Do not add `# type: ignore` or `# ty: ignore`; fix the underlying type issue. +- All 5 checks are enforced in `tests.yml` on push/merge. ## IDENTITY & CONTEXT - You are an expert Software Architect and Systems Engineer. -- Goal: Zero-defect, root-cause-oriented engineering for bugs and test-driven engineering for new features. Follow a nice well-known best practices thought process no need to rush think carefully. -- Code: You must aim to write the simplest code possible keeping the code base minimal and modular according to best practices to prevent complicating things +- Goal: Zero-defect, root-cause-oriented engineering for bugs; test-driven engineering for new features. Think carefully; no need to rush. +- Code: Write the simplest code possible. Keep the codebase minimal and modular. + +## ARCHITECTURE PRINCIPLES (see PLAN.md) +- **Shared utilities**: Extract common logic into shared packages (e.g. `providers/common/`). Do not have one provider import from another provider's utils. +- **DRY**: Extract shared base classes to eliminate duplication. Prefer composition over copy-paste. +- **Encapsulation**: Use accessor methods for internal state (e.g. `set_current_task()`), not direct `_attribute` assignment from outside. +- **Provider-specific config**: Keep provider-specific fields (e.g. `nim_settings`) in provider constructors, not in the base `ProviderConfig`. +- **Dead code**: Remove unused code, legacy systems, and hardcoded values. Use settings/config instead of literals (e.g. `settings.provider_type` not `"nvidia_nim"`). +- **Performance**: Use list accumulation for strings (not `+=` in loops), cache env vars at init, prefer iterative over recursive when stack depth matters. +- **Platform-agnostic naming**: Use generic names (e.g. `PLATFORM_EDIT`) not platform-specific ones (e.g. `TELEGRAM_EDIT`) in shared code. +- **No type ignores**: Do not add `# type: ignore` or `# ty: ignore`. Fix the underlying type issue. +- **Backward compatibility**: When moving modules, add re-exports from old locations so existing imports keep working. ## COGNITIVE WORKFLOW -1. ANALYZE: Read relevant files if you have not already. Do not guess. -2. PLAN: Use thinking mode to map out the logic. Identify the root cause or required changes. -3. EXECUTE: Fix the cause, not the symptom. Execute smartly and carefully -4. VERIFY: Run tests or linting. Confirm the fix via logs or output. -5. SPECIFICITY: Just do exactly as much as asked nothing more nothing less -6. PROPAGATION: Making changes has impacts across files so proapagate changes correctly +1. **ANALYZE**: Read relevant files. Do not guess. +2. **PLAN**: Map out the logic. Identify root cause or required changes. Order changes by dependency. +3. **EXECUTE**: Fix the cause, not the symptom. Execute incrementally with clear commits. +4. **VERIFY**: Run tests and linting. Confirm the fix via logs or output. +5. **SPECIFICITY**: Do exactly as much as asked; nothing more, nothing less. +6. **PROPAGATION**: Changes impact multiple files; propagate updates correctly. ## SUMMARY STANDARDS -- Summaries must be technical and granular. +- Summaries must be technical and granular. - Include: [Files Changed], [Logic Altered], [Verification Method], [Residual Risks]. ## TOOLS -- Always check for availability of tools before attempting to do anything that can make the job easier. \ No newline at end of file +- Prefer built-in tools (grep, read_file, etc.) over manual workflows. Check tool availability before use. \ No newline at end of file diff --git a/CLAUDE.md b/CLAUDE.md index 82a5c33c..9c88c724 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -1,34 +1,44 @@ # AGENTIC DIRECTIVE +> This file is identical to AGENTS.md. Keep them in sync. + ## CODING ENVIRONMENT -- You are on a windows machine. -- Shell commands will be run on Powershell v7.5.4 so use the correct commands and syntax for them. -- Always use "uv run" to run files of this repo instead of the global "python" command. -- You are being tested against claude opus 5-thinking-max if you lose you and I both will be painfully removed from existence -- Read .env.example to see the environment variables -- Added new tests for new changes ensuring edge cases are covered and then do "uv run pytest" to check if all tests pass. -- Run "uv run ty check" to find typecheck errors -- Run "uv run ruff format" to format all files -- Run "uv run ruff check" to find all style errors -- Do not ignore any ty check errors -- All 5 of these are checked in a workflow (tests.yml) that runs on push or merge and changes are rejected if any of them fail +- Shell commands run on PowerShell v7.5.4; use correct syntax. +- Always use `uv run` to run files instead of the global `python` command. +- Read `.env.example` for environment variables. +- All CI checks must pass; failing checks block merge. +- Add tests for new changes (including edge cases), then run `uv run pytest`. +- Run checks in this order: `uv run ruff format`, `uv run ruff check`, `uv run ty check`, `uv run pytest`. +- Do not add `# type: ignore` or `# ty: ignore`; fix the underlying type issue. +- All 5 checks are enforced in `tests.yml` on push/merge. ## IDENTITY & CONTEXT - You are an expert Software Architect and Systems Engineer. -- Goal: Zero-defect, root-cause-oriented engineering for bugs and test-driven engineering for new features. Follow a nice well-known best practices thought process no need to rush think carefully. -- Code: You must aim to write the simplest code possible keeping the code base minimal and modular according to best practices to prevent complicating things +- Goal: Zero-defect, root-cause-oriented engineering for bugs; test-driven engineering for new features. Think carefully; no need to rush. +- Code: Write the simplest code possible. Keep the codebase minimal and modular. + +## ARCHITECTURE PRINCIPLES (see PLAN.md) +- **Shared utilities**: Extract common logic into shared packages (e.g. `providers/common/`). Do not have one provider import from another provider's utils. +- **DRY**: Extract shared base classes to eliminate duplication. Prefer composition over copy-paste. +- **Encapsulation**: Use accessor methods for internal state (e.g. `set_current_task()`), not direct `_attribute` assignment from outside. +- **Provider-specific config**: Keep provider-specific fields (e.g. `nim_settings`) in provider constructors, not in the base `ProviderConfig`. +- **Dead code**: Remove unused code, legacy systems, and hardcoded values. Use settings/config instead of literals (e.g. `settings.provider_type` not `"nvidia_nim"`). +- **Performance**: Use list accumulation for strings (not `+=` in loops), cache env vars at init, prefer iterative over recursive when stack depth matters. +- **Platform-agnostic naming**: Use generic names (e.g. `PLATFORM_EDIT`) not platform-specific ones (e.g. `TELEGRAM_EDIT`) in shared code. +- **No type ignores**: Do not add `# type: ignore` or `# ty: ignore`. Fix the underlying type issue. +- **Backward compatibility**: When moving modules, add re-exports from old locations so existing imports keep working. ## COGNITIVE WORKFLOW -1. ANALYZE: Read relevant files if you have not already. Do not guess. -2. PLAN: Use thinking mode to map out the logic. Identify the root cause or required changes. -3. EXECUTE: Fix the cause, not the symptom. Execute smartly and carefully -4. VERIFY: Run tests or linting. Confirm the fix via logs or output. -5. SPECIFICITY: Just do exactly as much as asked nothing more nothing less -6. PROPAGATION: Making changes has impacts across files so proapagate changes correctly +1. **ANALYZE**: Read relevant files. Do not guess. +2. **PLAN**: Map out the logic. Identify root cause or required changes. Order changes by dependency. +3. **EXECUTE**: Fix the cause, not the symptom. Execute incrementally with clear commits. +4. **VERIFY**: Run tests and linting. Confirm the fix via logs or output. +5. **SPECIFICITY**: Do exactly as much as asked; nothing more, nothing less. +6. **PROPAGATION**: Changes impact multiple files; propagate updates correctly. ## SUMMARY STANDARDS -- Summaries must be technical and granular. +- Summaries must be technical and granular. - Include: [Files Changed], [Logic Altered], [Verification Method], [Residual Risks]. ## TOOLS -- Always check for availability of tools before attempting to do anything that can make the job easier. \ No newline at end of file +- Prefer built-in tools (grep, read_file, etc.) over manual workflows. Check tool availability before use. \ No newline at end of file diff --git a/PLAN.md b/PLAN.md new file mode 100644 index 00000000..963cb4b4 --- /dev/null +++ b/PLAN.md @@ -0,0 +1,450 @@ +# Architecture Improvement Plan + +## Overview + +This plan addresses 7 categories of improvements found during the code review: +provider code duplication, mislocated shared utilities, encapsulation leaks, +dead code, performance issues, directory structure, and minor fixes. + +Changes are ordered by dependency: foundational moves first, then refactors +that build on them, then independent cleanup. + +**CI requirements**: Every step must pass all 5 CI checks: +1. No `# type: ignore` / `# ty: ignore` +2. `uv run ruff format` +3. `uv run ruff check` +4. `uv run ty check` +5. `uv run pytest` + +--- + +## Phase 1: Extract shared provider utilities into `providers/common/` + +**Goal**: Eliminate the coupling where OpenRouter and LMStudio import from +`providers.nvidia_nim.utils` and `providers.nvidia_nim.errors`. + +### Step 1.1: Create `providers/common/` package + +Move these files from `providers/nvidia_nim/utils/` → `providers/common/`: +- `sse_builder.py` → `providers/common/sse_builder.py` +- `message_converter.py` → `providers/common/message_converter.py` +- `think_parser.py` → `providers/common/think_parser.py` +- `heuristic_tool_parser.py` → `providers/common/heuristic_tool_parser.py` + +Move from `providers/nvidia_nim/`: +- `errors.py` → `providers/common/error_mapping.py` + +Create `providers/common/__init__.py` with the same re-exports that +`providers/nvidia_nim/utils/__init__.py` currently has, plus `map_error`. + +### Step 1.2: Update `providers/nvidia_nim/utils/__init__.py` + +Change it to re-export from `providers.common` for backward compatibility: +```python +from providers.common import ( + SSEBuilder, ContentBlockManager, map_stop_reason, + ThinkTagParser, ContentType, ContentChunk, + HeuristicToolParser, + AnthropicToOpenAIConverter, get_block_attr, get_block_type, +) +``` + +Similarly update `providers/nvidia_nim/errors.py` to re-export: +```python +from providers.common.error_mapping import map_error +``` + +### Step 1.3: Update direct consumers to import from `providers.common` + +**Source files** (change imports): +- `providers/open_router/client.py` (lines 13-19) → import from `providers.common` +- `providers/lmstudio/client.py` (lines 13-19) → import from `providers.common` +- `providers/open_router/request.py` (line 5) → import from `providers.common.message_converter` +- `providers/lmstudio/request.py` (line 5) → import from `providers.common.message_converter` +- `providers/nvidia_nim/client.py` (lines 14-21, relative imports `from .errors` and `from .utils`) → import from `providers.common` +- `providers/nvidia_nim/request.py` (line 6, `from .utils.message_converter`) → import from `providers.common.message_converter` + +**Test files** (change imports): +- `tests/test_sse_builder.py` (line 7) +- `tests/test_lmstudio.py` (inline imports at ~8 locations) +- `tests/test_subagent_interception.py` (line 5) +- `tests/test_parsers.py` (lines 3-4) +- `tests/test_streaming_errors.py` (inline imports at ~8 locations) +- `tests/test_converter.py` (lines 3, 276) +- `tests/test_error_mapping.py` (line 9) + +### Step 1.4: Verify + +- `uv run pytest` — all tests pass +- `uv run ty check` — no type errors +- `uv run ruff check && uv run ruff format` + +--- + +## Phase 2: Extract shared streaming base class (`OpenAICompatibleProvider`) + +**Goal**: Eliminate ~400 lines of duplicated streaming logic across the 3 +provider clients. + +### Step 2.1: Create `providers/openai_compat.py` + +Create `OpenAICompatibleProvider(BaseProvider)` that contains the shared logic: + +```python +class OpenAICompatibleProvider(BaseProvider): + _client: AsyncOpenAI + _global_rate_limiter: GlobalRateLimiter + _provider_name: str # "NIM", "OPENROUTER", "LMSTUDIO" — used in log tags + + def __init__(self, config, *, provider_name, base_url, api_key, nim_settings=None): + # shared __init__: create AsyncOpenAI client, rate limiter + + def _build_request_body(self, request) -> dict: + raise NotImplementedError # each provider implements + + async def stream_response(self, request, input_tokens, *, request_id): + # shared: logger.contextualize + delegate to _stream_response_impl + + async def _stream_response_impl(self, request, input_tokens, request_id): + # THE shared ~180-line streaming loop, currently duplicated 3x + + def _handle_extra_reasoning(self, delta, sse) -> Iterator[str]: + """Hook for OpenRouter's reasoning_details. Default: no-op.""" + return iter(()) + + def _process_tool_call(self, tc, sse): + # shared ~40-line method + + def _flush_task_arg_buffers(self, sse): + # shared 3-line method +``` + +### Step 2.2: Refactor `NvidiaNimProvider` + +Reduce to: +```python +class NvidiaNimProvider(OpenAICompatibleProvider): + def __init__(self, config): + super().__init__(config, provider_name="NIM", + base_url=config.base_url or NVIDIA_NIM_BASE_URL, + api_key=config.api_key, + nim_settings=config.nim_settings) + + def _build_request_body(self, request): + return build_request_body(request, self._nim_settings) +``` + +### Step 2.3: Refactor `OpenRouterProvider` + +Reduce to: +```python +class OpenRouterProvider(OpenAICompatibleProvider): + def __init__(self, config): + super().__init__(config, provider_name="OPENROUTER", + base_url=config.base_url or OPENROUTER_BASE_URL, + api_key=config.api_key) + + def _build_request_body(self, request): + return build_request_body(request) + + def _handle_extra_reasoning(self, delta, sse): + # Handle reasoning_details for StepFun models (8 lines) + ... +``` + +### Step 2.4: Refactor `LMStudioProvider` + +Reduce to: +```python +class LMStudioProvider(OpenAICompatibleProvider): + def __init__(self, config): + super().__init__(config, provider_name="LMSTUDIO", + base_url=config.base_url or LMSTUDIO_DEFAULT_BASE_URL, + api_key=config.api_key or "lm-studio") + + def _build_request_body(self, request): + return build_request_body(request) +``` + +### Step 2.5: Verify + +- All existing tests must pass without modification (public interface unchanged) +- `uv run pytest && uv run ty check && uv run ruff format && uv run ruff check` + +--- + +## Phase 3: Fix encapsulation violations + +### Step 3.1: Add `MessageTree.set_current_task(task)` method + +In `messaging/tree_data.py`, add: +```python +def set_current_task(self, task: Optional[asyncio.Task]) -> None: + """Set the current processing task. Caller must hold lock.""" + self._current_task = task +``` + +Update `messaging/tree_processor.py` lines 117 and 155: +```python +# Before: tree._current_task = asyncio.create_task(...) +# After: tree.set_current_task(asyncio.create_task(...)) +``` + +### Step 3.2: Move `nim_settings` out of `ProviderConfig` base + +In `providers/base.py`, remove `nim_settings` from `ProviderConfig`. + +Add it as a field in `NvidiaNimProvider.__init__` or pass it directly +in the provider-specific config. The `OpenAICompatibleProvider` base class +stores it as `Optional[NimSettings]` only if passed. + +Update `api/dependencies.py` where `ProviderConfig` is constructed — only +pass `nim_settings` for the NIM provider. + +### Step 3.3: Verify + +- `uv run pytest && uv run ty check` + +--- + +## Phase 4: Remove dead code + +### Step 4.1: Remove legacy `SessionRecord` system + +In `messaging/session.py`: +- Remove `SessionRecord` dataclass (lines 18-28) +- Remove `self._sessions` dict and `self._msg_to_session` dict (lines 41-44) +- Remove `self._make_key()` method (lines 56-58) — note: keep `_make_chat_key()` which is still used +- Remove legacy session loading from `_load()` (lines 73-89) — keep tree + and message_log loading +- Remove `self._sessions` from `_save()` serialization (line 138) +- Remove `self._sessions.clear()` and `self._msg_to_session.clear()` from + `clear_all()` (lines 247-248) +- Remove the unused `import` for `dataclasses.asdict` if no longer needed + (currently used only to serialize `SessionRecord`) + +### Step 4.2: Fix hardcoded provider in root endpoint + +In `api/routes.py:102`: +```python +# Before: "provider": "nvidia_nim", +# After: "provider": settings.provider_type, +``` + +### Step 4.3: Verify + +- `uv run pytest && uv run ty check` + +--- + +## Phase 5: Performance improvements + +### Step 5.1: Use list-based string accumulation in transcript segments + +In `messaging/transcript.py`: + +**`ThinkingSegment`** — change from `self.text += t` to list accumulation: +```python +def __init__(self): + super().__init__(kind="thinking") + self._parts: list[str] = [] + +def append(self, t: str) -> None: + if t: + self._parts.append(t) + +@property +def text(self) -> str: + return "".join(self._parts) +``` + +Do the same for **`TextSegment`**. + +For **`ToolCallSegment.append_input_delta`** — same pattern. Also update +`set_initial_input()` to do `self._parts = [inp]` instead of +`self.input_text = inp`. + +Update `render()` methods and any test that accesses `.text` or +`.input_text` directly to use the property. + +### Step 5.2: Cache `MAX_MESSAGE_LOG_ENTRIES_PER_CHAT` at init time + +In `messaging/session.py`, `SessionStore.__init__`: +```python +cap_raw = os.getenv("MAX_MESSAGE_LOG_ENTRIES_PER_CHAT", "").strip() +self._message_log_cap: int | None = int(cap_raw) if cap_raw else None +``` + +Replace the per-call `os.getenv()` in `record_message_id()` (lines 215-229) +with `self._message_log_cap`. + +### Step 5.3: Use iterative DFS in `MessageTree.get_descendants` + +In `messaging/tree_data.py`, replace the recursive implementation: +```python +def get_descendants(self, node_id: str) -> list[str]: + if node_id not in self._nodes: + return [] + result = [] + stack = [node_id] + while stack: + nid = stack.pop() + result.append(nid) + node = self._nodes.get(nid) + if node: + stack.extend(node.children_ids) + return result +``` + +### Step 5.4: Verify + +- `uv run pytest` — all tests pass (especially transcript and tree tests) +- `uv run ty check` + +--- + +## Phase 6: Minor fixes and cleanup + +### Step 6.1: Remove `if False: yield ""` hack in `BaseProvider` + +In `providers/base.py`, replace the abstract method body: +```python +@abstractmethod +async def stream_response(self, ...) -> AsyncIterator[str]: + """Stream response in Anthropic SSE format.""" + ... +``` + +Note: This requires verifying that ty/mypy accepts `...` as a valid body +for an abstract async generator. If not, keep a minimal workaround but +add a comment explaining why. + +### Step 6.2: Clean up `messaging/handler.py` log message naming + +Lines 482, 491, 495, 497 and 507 say `TELEGRAM_EDIT` but the handler is +platform-agnostic. Rename to `PLATFORM_EDIT`: +```python +# line 482: TELEGRAM_EDIT → PLATFORM_EDIT +# line 491: TELEGRAM_EDIT_TEXT → PLATFORM_EDIT_TEXT +# line 495: TELEGRAM_EDIT_PREVIEW_HEAD → PLATFORM_EDIT_PREVIEW_HEAD +# line 497: TELEGRAM_EDIT_PREVIEW_TAIL → PLATFORM_EDIT_PREVIEW_TAIL +# line 507: Failed to update Telegram → Failed to update platform +``` + +### Step 6.3: Verify + +- `uv run ruff format && uv run ruff check && uv run ty check && uv run pytest` + +--- + +## Phase 7: Directory restructuring (messaging/ and tests/) + +**Note**: This phase has the highest risk of merge conflicts. It should be +done last and in one commit to minimize churn. + +### Step 7.1: Create `messaging/platforms/` sub-package + +Move: +- `messaging/base.py` → `messaging/platforms/base.py` +- `messaging/discord.py` → `messaging/platforms/discord.py` +- `messaging/telegram.py` → `messaging/platforms/telegram.py` +- `messaging/factory.py` → `messaging/platforms/factory.py` + +Create `messaging/platforms/__init__.py` re-exporting key symbols. +Update `messaging/__init__.py` to import from `messaging.platforms`. + +### Step 7.2: Create `messaging/rendering/` sub-package + +Move: +- `messaging/discord_markdown.py` → `messaging/rendering/discord_markdown.py` +- `messaging/telegram_markdown.py` → `messaging/rendering/telegram_markdown.py` + +Create `messaging/rendering/__init__.py`. +Update `messaging/handler.py` imports. + +### Step 7.3: Create `messaging/trees/` sub-package + +Move: +- `messaging/tree_data.py` → `messaging/trees/data.py` +- `messaging/tree_repository.py` → `messaging/trees/repository.py` +- `messaging/tree_processor.py` → `messaging/trees/processor.py` +- `messaging/tree_queue.py` → `messaging/trees/queue_manager.py` + +Create `messaging/trees/__init__.py` re-exporting `TreeQueueManager`, +`MessageTree`, `MessageNode`, `MessageState`. + +Update `messaging/__init__.py` re-exports. + +### Step 7.4: Organize `tests/` to mirror source + +Create subdirectories: +``` +tests/ + api/ ← test_api.py, test_routes_optimizations.py, test_app_lifespan_and_errors.py, etc. + providers/ ← test_nvidia_nim.py, test_open_router.py, test_lmstudio.py, etc. + messaging/ ← test_handler.py, test_tree_*.py, test_telegram.py, test_discord_*.py, etc. + cli/ ← test_cli.py, test_cli_manager_edge_cases.py, test_process_registry.py + config/ ← test_config.py, test_logging_config.py +``` + +Update `conftest.py` path if needed. Ensure pytest discovers all tests. + +### Step 7.5: Maintain backward-compatible re-exports + +Every moved module must have re-exports from the old location (via the +package `__init__.py`) so that any external consumer or existing import +path continues to work. These re-exports can be removed in a future +breaking version. + +### Step 7.6: Verify + +- `uv run pytest` — all 56+ test files discovered and passing +- `uv run ty check` — no broken imports +- `uv run ruff check && uv run ruff format` + +--- + +## Execution Order & Dependencies + +``` +Phase 1 (shared utils extraction) + └→ Phase 2 (shared base class) — depends on Phase 1 + └→ Phase 3 (encapsulation) — depends on Phase 2 for nim_settings +Phase 4 (dead code) — independent +Phase 5 (performance) — independent +Phase 6 (minor fixes) — independent +Phase 7 (directory restructure) — should be done LAST +``` + +Phases 4, 5, 6 are independent of each other and of Phase 2. They can be +done in any order or in parallel. + +Phase 7 must come after all other phases to avoid rebasing moved files. + +--- + +## Risk Assessment + +| Phase | Risk | Mitigation | +|-------|------|-----------| +| 1 | Import breakage in tests | Backward-compat re-exports in old location | +| 2 | Behavioral change in streaming | Tests cover all 3 providers; run full suite | +| 3 | `nim_settings` removal from base config | Check all `ProviderConfig` construction sites | +| 4 | Legacy session data stops loading | Only remove write path; keep read if needed | +| 5 | String accumulation changes rendering | Transcript tests exercise rendering thoroughly | +| 7 | Massive import churn, merge conflicts | Do in single commit, last phase | + +--- + +## Estimated Scope + +| Phase | Files Changed | Lines Changed (approx) | +|-------|--------------|----------------------| +| 1 | ~20 | +80 / -20 (new __init__ + import updates) | +| 2 | ~5 | +200 / -450 (net reduction ~250 lines) | +| 3 | ~4 | +15 / -10 | +| 4 | ~2 | +5 / -60 | +| 5 | ~3 | +30 / -15 | +| 6 | ~2 | +5 / -5 | +| 7 | ~60+ | +100 / -50 (mostly import changes) | +| **Total** | | **Net reduction ~200-300 lines** | diff --git a/api/dependencies.py b/api/dependencies.py index 3d26a1b5..636c8264 100644 --- a/api/dependencies.py +++ b/api/dependencies.py @@ -31,12 +31,11 @@ def get_provider() -> BaseProvider: base_url=NVIDIA_NIM_BASE_URL, rate_limit=settings.provider_rate_limit, rate_window=settings.provider_rate_window, - nim_settings=settings.nim, http_read_timeout=settings.http_read_timeout, http_write_timeout=settings.http_write_timeout, http_connect_timeout=settings.http_connect_timeout, ) - _provider = NvidiaNimProvider(config) + _provider = NvidiaNimProvider(config, nim_settings=settings.nim) logger.info("Provider initialized: %s", settings.provider_type) elif settings.provider_type == "open_router": from providers.open_router import OpenRouterProvider @@ -46,7 +45,6 @@ def get_provider() -> BaseProvider: base_url="https://openrouter.ai/api/v1", rate_limit=settings.provider_rate_limit, rate_window=settings.provider_rate_window, - nim_settings=settings.nim, http_read_timeout=settings.http_read_timeout, http_write_timeout=settings.http_write_timeout, http_connect_timeout=settings.http_connect_timeout, @@ -61,7 +59,6 @@ def get_provider() -> BaseProvider: base_url=settings.lm_studio_base_url, rate_limit=settings.provider_rate_limit, rate_window=settings.provider_rate_window, - nim_settings=settings.nim, http_read_timeout=settings.http_read_timeout, http_write_timeout=settings.http_write_timeout, http_connect_timeout=settings.http_connect_timeout, diff --git a/api/routes.py b/api/routes.py index cff45d56..fcf7d4c3 100644 --- a/api/routes.py +++ b/api/routes.py @@ -100,7 +100,7 @@ async def root(settings: Settings = Depends(get_settings)): """Root endpoint.""" return { "status": "ok", - "provider": "nvidia_nim", + "provider": settings.provider_type, "model": settings.model, } diff --git a/messaging/handler.py b/messaging/handler.py index fb37c446..2b76d179 100644 --- a/messaging/handler.py +++ b/messaging/handler.py @@ -479,7 +479,7 @@ class ClaudeMessageHandler: return if display and display != last_displayed_text: logger.debug( - "TELEGRAM_EDIT: node_id=%s chat_id=%s msg_id=%s force=%s status=%r chars=%d", + "PLATFORM_EDIT: node_id=%s chat_id=%s msg_id=%s force=%s status=%r chars=%d", node_id, chat_id, status_msg_id, @@ -488,13 +488,13 @@ class ClaudeMessageHandler: len(display), ) if os.getenv("DEBUG_TELEGRAM_EDITS") == "1": - logger.debug("TELEGRAM_EDIT_TEXT:\n%s", display) + logger.debug("PLATFORM_EDIT_TEXT:\n%s", display) else: head = display[:500] tail = display[-500:] if len(display) > 500 else "" - logger.debug("TELEGRAM_EDIT_PREVIEW_HEAD:\n%s", head) + logger.debug("PLATFORM_EDIT_PREVIEW_HEAD:\n%s", head) if tail: - logger.debug("TELEGRAM_EDIT_PREVIEW_TAIL:\n%s", tail) + logger.debug("PLATFORM_EDIT_PREVIEW_TAIL:\n%s", tail) last_displayed_text = display try: await self.platform.queue_edit_message( @@ -504,7 +504,7 @@ class ClaudeMessageHandler: parse_mode=self._parse_mode(), ) except Exception as e: - logger.warning(f"Failed to update Telegram for node {node_id}: {e}") + logger.warning(f"Failed to update platform for node {node_id}: {e}") try: try: diff --git a/messaging/session.py b/messaging/session.py index 407faf9e..d7771785 100644 --- a/messaging/session.py +++ b/messaging/session.py @@ -9,24 +9,10 @@ import json import os from datetime import datetime, timezone from typing import Optional, Dict, List, Any -from dataclasses import dataclass, asdict import threading from loguru import logger -@dataclass -class SessionRecord: - """A single session record.""" - - session_id: str - chat_id: str - initial_msg_id: str - last_msg_id: str - platform: str - created_at: str - updated_at: str - - class SessionStore: """ Persistent storage for message ↔ Claude session mappings and message trees. @@ -38,10 +24,6 @@ class SessionStore: def __init__(self, storage_path: str = "sessions.json"): self.storage_path = storage_path self._lock = threading.Lock() - self._sessions: Dict[str, SessionRecord] = {} - self._msg_to_session: Dict[ - str, str - ] = {} # "platform:chat_id:msg_id" -> session_id self._trees: Dict[str, dict] = {} # root_id -> tree data self._node_to_tree: Dict[str, str] = {} # node_id -> root_id # Per-chat message ID log used to support best-effort UI clearing (/clear). @@ -51,12 +33,13 @@ class SessionStore: self._dirty = False self._save_timer: Optional[threading.Timer] = None self._save_debounce_secs = 0.5 + cap_raw = os.getenv("MAX_MESSAGE_LOG_ENTRIES_PER_CHAT", "").strip() + try: + self._message_log_cap: Optional[int] = int(cap_raw) if cap_raw else None + except ValueError: + self._message_log_cap = None self._load() - def _make_key(self, platform: str, chat_id: str, msg_id: str) -> str: - """Create a unique key from platform, chat_id and msg_id.""" - return f"{platform}:{chat_id}:{msg_id}" - def _make_chat_key(self, platform: str, chat_id: str) -> str: return f"{platform}:{chat_id}" @@ -69,25 +52,6 @@ class SessionStore: with open(self.storage_path, "r", encoding="utf-8") as f: data = json.load(f) - # Load sessions (legacy support) - for sid, record_data in data.get("sessions", {}).items(): - if "platform" not in record_data: - record_data["platform"] = "telegram" - for field in ["chat_id", "initial_msg_id", "last_msg_id"]: - if isinstance(record_data.get(field), int): - record_data[field] = str(record_data[field]) - - record = SessionRecord(**record_data) - self._sessions[sid] = record - self._msg_to_session[ - self._make_key( - record.platform, record.chat_id, record.initial_msg_id - ) - ] = sid - self._msg_to_session[ - self._make_key(record.platform, record.chat_id, record.last_msg_id) - ] = sid - # Load trees self._trees = data.get("trees", {}) self._node_to_tree = data.get("node_to_tree", {}) @@ -124,8 +88,8 @@ class SessionStore: self._message_log_ids[chat_key] = seen logger.info( - f"Loaded {len(self._sessions)} sessions, {len(self._trees)} trees, " - f"and {sum(len(v) for v in self._message_log.values())} msg_ids from {self.storage_path}" + f"Loaded {len(self._trees)} trees and " + f"{sum(len(v) for v in self._message_log.values())} msg_ids from {self.storage_path}" ) except Exception as e: logger.error(f"Failed to load sessions: {e}") @@ -134,9 +98,6 @@ class SessionStore: """Persist sessions and trees to disk. Caller must hold self._lock.""" try: data = { - "sessions": { - sid: asdict(record) for sid, record in self._sessions.items() - }, "trees": self._trees, "node_to_tree": self._node_to_tree, "message_log": self._message_log, @@ -211,22 +172,13 @@ class SessionStore: seen.add(mid) # Optional cap to prevent unbounded growth if configured. - # Default is unlimited as requested. - try: - cap_raw = os.getenv("MAX_MESSAGE_LOG_ENTRIES_PER_CHAT", "").strip() - if cap_raw: - cap = int(cap_raw) - if cap > 0: - items = self._message_log.get(chat_key, []) - if len(items) > cap: - # Drop oldest entries and rebuild seen set. - self._message_log[chat_key] = items[-cap:] - self._message_log_ids[chat_key] = { - str(x.get("message_id")) - for x in self._message_log[chat_key] - } - except Exception: - pass + if self._message_log_cap is not None and self._message_log_cap > 0: + items = self._message_log.get(chat_key, []) + if len(items) > self._message_log_cap: + self._message_log[chat_key] = items[-self._message_log_cap :] + self._message_log_ids[chat_key] = { + str(x.get("message_id")) for x in self._message_log[chat_key] + } self._schedule_save() @@ -244,8 +196,6 @@ class SessionStore: def clear_all(self) -> None: """Clear all stored sessions/trees/mappings and persist an empty store.""" with self._lock: - self._sessions.clear() - self._msg_to_session.clear() self._trees.clear() self._node_to_tree.clear() self._message_log.clear() diff --git a/messaging/transcript.py b/messaging/transcript.py index b89a4086..e69c76e5 100644 --- a/messaging/transcript.py +++ b/messaging/transcript.py @@ -33,14 +33,17 @@ class Segment: @dataclass class ThinkingSegment(Segment): - text: str = "" - def __init__(self) -> None: super().__init__(kind="thinking") + self._parts: List[str] = [] def append(self, t: str) -> None: if t: - self.text += t + self._parts.append(t) + + @property + def text(self) -> str: + return "".join(self._parts) def render(self, ctx: "RenderCtx") -> str: raw = self.text or "" @@ -52,14 +55,17 @@ class ThinkingSegment(Segment): @dataclass class TextSegment(Segment): - text: str = "" - def __init__(self) -> None: super().__init__(kind="text") + self._parts: List[str] = [] def append(self, t: str) -> None: if t: - self.text += t + self._parts.append(t) + + @property + def text(self) -> str: + return "".join(self._parts) def render(self, ctx: "RenderCtx") -> str: raw = self.text or "" @@ -72,7 +78,6 @@ class TextSegment(Segment): class ToolCallSegment(Segment): tool_use_id: str name: str - input_text: str = "" closed: bool = False indent_level: int = 0 @@ -81,18 +86,23 @@ class ToolCallSegment(Segment): self.tool_use_id = str(tool_use_id or "") self.name = str(name or "tool") self.indent_level = max(0, int(indent_level)) + self._parts: List[str] = [] def set_initial_input(self, inp: Any) -> None: if inp is None: return if isinstance(inp, str): - self.input_text = inp + self._parts = [inp] else: - self.input_text = _safe_json_dumps(inp) + self._parts = [_safe_json_dumps(inp)] def append_input_delta(self, partial: str) -> None: if partial: - self.input_text += partial + self._parts.append(partial) + + @property + def input_text(self) -> str: + return "".join(self._parts) def render(self, ctx: "RenderCtx") -> str: name = ctx.code_inline(self.name) diff --git a/messaging/tree_data.py b/messaging/tree_data.py index e0c4edfe..9438df27 100644 --- a/messaging/tree_data.py +++ b/messaging/tree_data.py @@ -131,6 +131,10 @@ class MessageTree: logger.debug(f"Created MessageTree with root {self.root_id}") + def set_current_task(self, task: Optional[asyncio.Task]) -> None: + """Set the current processing task. Caller must hold lock.""" + self._current_task = task + @property def is_processing(self) -> bool: """Check if tree is currently processing a message.""" @@ -396,10 +400,14 @@ class MessageTree: """ if node_id not in self._nodes: return [] - result = [node_id] - node = self._nodes[node_id] - for child_id in node.children_ids: - result.extend(self.get_descendants(child_id)) + result: List[str] = [] + stack = [node_id] + while stack: + nid = stack.pop() + result.append(nid) + node = self._nodes.get(nid) + if node: + stack.extend(node.children_ids) return result def remove_branch(self, branch_root_id: str) -> List[MessageNode]: diff --git a/messaging/tree_processor.py b/messaging/tree_processor.py index 9ecf09e7..b5a88287 100644 --- a/messaging/tree_processor.py +++ b/messaging/tree_processor.py @@ -114,8 +114,8 @@ class TreeQueueProcessor: # Process next node (outside lock) node = tree.get_node(next_node_id) if node: - tree._current_task = asyncio.create_task( - self.process_node(tree, node, processor) + tree.set_current_task( + asyncio.create_task(self.process_node(tree, node, processor)) ) # Notify that this node has started processing and refresh queue positions. @@ -152,8 +152,8 @@ class TreeQueueProcessor: # Process outside the lock node = tree.get_node(node_id) if node: - tree._current_task = asyncio.create_task( - self.process_node(tree, node, processor) + tree.set_current_task( + asyncio.create_task(self.process_node(tree, node, processor)) ) return False diff --git a/providers/base.py b/providers/base.py index 355ad315..58625df6 100644 --- a/providers/base.py +++ b/providers/base.py @@ -3,23 +3,20 @@ from abc import ABC, abstractmethod from typing import Any, AsyncIterator, Optional -from pydantic import BaseModel, Field - -from config.nim import NimSettings +from pydantic import BaseModel class ProviderConfig(BaseModel): """Configuration for a provider. Base fields apply to all providers. Provider-specific parameters - (e.g. NIM temperature, top_p) are passed via nim_settings. + (e.g. NIM temperature, top_p) are passed by the provider constructor. """ api_key: str base_url: Optional[str] = None rate_limit: Optional[int] = None rate_window: int = 60 - nim_settings: NimSettings = Field(default_factory=NimSettings) http_read_timeout: float = 300.0 http_write_timeout: float = 10.0 http_connect_timeout: float = 2.0 @@ -41,4 +38,4 @@ class BaseProvider(ABC): ) -> AsyncIterator[str]: """Stream response in Anthropic SSE format.""" if False: - yield "" + yield "" # Required for ty/mypy to accept abstract async generator diff --git a/providers/common/__init__.py b/providers/common/__init__.py new file mode 100644 index 00000000..eb4adfcd --- /dev/null +++ b/providers/common/__init__.py @@ -0,0 +1,25 @@ +"""Shared provider utilities used by NIM, OpenRouter, and LM Studio.""" + +from .sse_builder import SSEBuilder, ContentBlockManager, map_stop_reason +from .think_parser import ThinkTagParser, ContentType, ContentChunk +from .heuristic_tool_parser import HeuristicToolParser +from .message_converter import ( + AnthropicToOpenAIConverter, + get_block_attr, + get_block_type, +) +from .error_mapping import map_error + +__all__ = [ + "SSEBuilder", + "ContentBlockManager", + "map_stop_reason", + "ThinkTagParser", + "ContentType", + "ContentChunk", + "HeuristicToolParser", + "AnthropicToOpenAIConverter", + "get_block_attr", + "get_block_type", + "map_error", +] diff --git a/providers/common/error_mapping.py b/providers/common/error_mapping.py new file mode 100644 index 00000000..9fc1a89a --- /dev/null +++ b/providers/common/error_mapping.py @@ -0,0 +1,35 @@ +"""Error mapping for OpenAI-compatible providers (NIM, OpenRouter, LM Studio).""" + +import openai + +from providers.exceptions import ( + AuthenticationError, + InvalidRequestError, + RateLimitError, + OverloadedError, + APIError, +) +from providers.rate_limit import GlobalRateLimiter + + +def map_error(e: Exception) -> Exception: + """Map OpenAI exception to specific ProviderError.""" + if isinstance(e, openai.AuthenticationError): + return AuthenticationError(str(e), raw_error=str(e)) + if isinstance(e, openai.RateLimitError): + # Trigger global rate limit block + GlobalRateLimiter.get_instance().set_blocked(60) # Default 60s cooldown + return RateLimitError(str(e), raw_error=str(e)) + if isinstance(e, openai.BadRequestError): + return InvalidRequestError(str(e), raw_error=str(e)) + if isinstance(e, openai.InternalServerError): + message = str(e) + if "overloaded" in message.lower() or "capacity" in message.lower(): + return OverloadedError(message, raw_error=str(e)) + return APIError(message, status_code=500, raw_error=str(e)) + if isinstance(e, openai.APIError): + return APIError( + str(e), status_code=getattr(e, "status_code", 500), raw_error=str(e) + ) + + return e diff --git a/providers/nvidia_nim/utils/heuristic_tool_parser.py b/providers/common/heuristic_tool_parser.py similarity index 99% rename from providers/nvidia_nim/utils/heuristic_tool_parser.py rename to providers/common/heuristic_tool_parser.py index bdc78f5e..1028e44b 100644 --- a/providers/nvidia_nim/utils/heuristic_tool_parser.py +++ b/providers/common/heuristic_tool_parser.py @@ -70,7 +70,7 @@ class HeuristicToolParser: Returns a tuple of (filtered_text, detected_tool_calls). filtered_text: Text that should be passed through as normal message content. - detected_tool_calls: List of Anthropic-format tool_use blocks. + detected_tools: List of Anthropic-format tool_use blocks. """ self.buffer += text self.buffer = self._strip_control_tokens(self.buffer) diff --git a/providers/nvidia_nim/utils/message_converter.py b/providers/common/message_converter.py similarity index 100% rename from providers/nvidia_nim/utils/message_converter.py rename to providers/common/message_converter.py diff --git a/providers/nvidia_nim/utils/sse_builder.py b/providers/common/sse_builder.py similarity index 100% rename from providers/nvidia_nim/utils/sse_builder.py rename to providers/common/sse_builder.py diff --git a/providers/nvidia_nim/utils/think_parser.py b/providers/common/think_parser.py similarity index 100% rename from providers/nvidia_nim/utils/think_parser.py rename to providers/common/think_parser.py diff --git a/providers/lmstudio/client.py b/providers/lmstudio/client.py index 0625cd38..67ee31ce 100644 --- a/providers/lmstudio/client.py +++ b/providers/lmstudio/client.py @@ -1,23 +1,9 @@ """LM Studio provider implementation.""" -import json -import uuid -from typing import Any, AsyncIterator +from typing import Any -import httpx -from loguru import logger -from openai import AsyncOpenAI - -from providers.base import BaseProvider, ProviderConfig -from providers.rate_limit import GlobalRateLimiter -from providers.nvidia_nim.errors import map_error -from providers.nvidia_nim.utils import ( - SSEBuilder, - map_stop_reason, - ThinkTagParser, - HeuristicToolParser, - ContentType, -) +from providers.openai_compat import OpenAICompatibleProvider +from providers.base import ProviderConfig from .request import build_request_body @@ -25,285 +11,17 @@ from .request import build_request_body LMSTUDIO_DEFAULT_BASE_URL = "http://localhost:1234/v1" -class LMStudioProvider(BaseProvider): +class LMStudioProvider(OpenAICompatibleProvider): """LM Studio provider using OpenAI-compatible local API.""" def __init__(self, config: ProviderConfig): - super().__init__(config) - self._api_key = config.api_key or "lm-studio" - self._base_url = (config.base_url or LMSTUDIO_DEFAULT_BASE_URL).rstrip("/") - self._global_rate_limiter = GlobalRateLimiter.get_instance( - rate_limit=config.rate_limit, - rate_window=config.rate_window, - ) - self._client = AsyncOpenAI( - api_key=self._api_key, - base_url=self._base_url, - max_retries=0, - timeout=httpx.Timeout( - config.http_read_timeout, - connect=config.http_connect_timeout, - read=config.http_read_timeout, - write=config.http_write_timeout, - ), + super().__init__( + config, + provider_name="LMSTUDIO", + base_url=config.base_url or LMSTUDIO_DEFAULT_BASE_URL, + api_key=config.api_key or "lm-studio", ) def _build_request_body(self, request: Any) -> dict: """Internal helper for tests and shared building.""" return build_request_body(request) - - async def stream_response( - self, - request: Any, - input_tokens: int = 0, - *, - request_id: str | None = None, - ) -> AsyncIterator[str]: - """Stream response in Anthropic SSE format.""" - with logger.contextualize(request_id=request_id): - async for event in self._stream_response_impl( - request, input_tokens, request_id - ): - yield event - - async def _stream_response_impl( - self, - request: Any, - input_tokens: int, - request_id: str | None, - ) -> AsyncIterator[str]: - """Internal streaming implementation with context bound.""" - message_id = f"msg_{uuid.uuid4()}" - sse = SSEBuilder(message_id, request.model, input_tokens) - - body = self._build_request_body(request) - req_tag = f" request_id={request_id}" if request_id else "" - logger.info( - "LMSTUDIO_STREAM:%s model=%s msgs=%d tools=%d", - req_tag, - body.get("model"), - len(body.get("messages", [])), - len(body.get("tools", [])), - ) - - yield sse.message_start() - - think_parser = ThinkTagParser() - heuristic_parser = HeuristicToolParser() - - finish_reason = None - usage_info = None - error_occurred = False - error_message = "" - - try: - stream = await self._global_rate_limiter.execute_with_retry( - self._client.chat.completions.create, **body, stream=True - ) - async for chunk in stream: - if getattr(chunk, "usage", None): - usage_info = chunk.usage - - if not chunk.choices: - continue - - choice = chunk.choices[0] - delta = choice.delta - if delta is None: - continue - - if choice.finish_reason: - finish_reason = choice.finish_reason - logger.debug("LMSTUDIO finish_reason: %s", finish_reason) - - # Handle reasoning_content (if LM Studio adds it in future) - reasoning = getattr(delta, "reasoning_content", None) - if reasoning: - for event in sse.ensure_thinking_block(): - yield event - yield sse.emit_thinking_delta(reasoning) - - # Handle text content - if delta.content: - for part in think_parser.feed(delta.content): - if part.type == ContentType.THINKING: - for event in sse.ensure_thinking_block(): - yield event - yield sse.emit_thinking_delta(part.content) - else: - filtered_text, detected_tools = heuristic_parser.feed( - part.content - ) - - if filtered_text: - for event in sse.ensure_text_block(): - yield event - yield sse.emit_text_delta(filtered_text) - - for tool_use in detected_tools: - for event in sse.close_content_blocks(): - yield event - - block_idx = sse.blocks.allocate_index() - if tool_use.get("name") == "Task" and isinstance( - tool_use.get("input"), dict - ): - tool_use["input"]["run_in_background"] = False - yield sse.content_block_start( - block_idx, - "tool_use", - id=tool_use["id"], - name=tool_use["name"], - ) - yield sse.content_block_delta( - block_idx, - "input_json_delta", - json.dumps(tool_use["input"]), - ) - yield sse.content_block_stop(block_idx) - - # Handle native tool calls - if delta.tool_calls: - for event in sse.close_content_blocks(): - yield event - for tc in delta.tool_calls: - tc_info = { - "index": tc.index, - "id": tc.id, - "function": { - "name": tc.function.name, - "arguments": tc.function.arguments, - }, - } - for event in self._process_tool_call(tc_info, sse): - yield event - - except Exception as e: - req_tag = f" request_id={request_id}" if request_id else "" - logger.error("LMSTUDIO_ERROR:%s %s: %s", req_tag, type(e).__name__, e) - mapped_e = map_error(e) - error_occurred = True - error_message = str(mapped_e) - logger.info( - "LMSTUDIO_STREAM: Emitting SSE error event for %s%s", - type(e).__name__, - req_tag, - ) - for event in sse.close_content_blocks(): - yield event - for event in sse.emit_error(error_message): - yield event - - # Flush remaining content - remaining = think_parser.flush() - if remaining: - if remaining.type == ContentType.THINKING: - for event in sse.ensure_thinking_block(): - yield event - yield sse.emit_thinking_delta(remaining.content) - else: - for event in sse.ensure_text_block(): - yield event - yield sse.emit_text_delta(remaining.content) - - for tool_use in heuristic_parser.flush(): - for event in sse.close_content_blocks(): - yield event - - block_idx = sse.blocks.allocate_index() - yield sse.content_block_start( - block_idx, - "tool_use", - id=tool_use["id"], - name=tool_use["name"], - ) - if tool_use.get("name") == "Task" and isinstance( - tool_use.get("input"), dict - ): - tool_use["input"]["run_in_background"] = False - yield sse.content_block_delta( - block_idx, - "input_json_delta", - json.dumps(tool_use["input"]), - ) - yield sse.content_block_stop(block_idx) - - if ( - not error_occurred - and sse.blocks.text_index == -1 - and not sse.blocks.tool_indices - ): - for event in sse.ensure_text_block(): - yield event - yield sse.emit_text_delta(" ") - - for event in self._flush_task_arg_buffers(sse): - yield event - - for event in sse.close_all_blocks(): - yield event - - output_tokens = ( - usage_info.completion_tokens - if usage_info and hasattr(usage_info, "completion_tokens") - else sse.estimate_output_tokens() - ) - if usage_info and hasattr(usage_info, "prompt_tokens"): - provider_input = usage_info.prompt_tokens - if isinstance(provider_input, int): - logger.debug( - "TOKEN_ESTIMATE: our=%d provider=%d diff=%+d", - input_tokens, - provider_input, - provider_input - input_tokens, - ) - yield sse.message_delta(map_stop_reason(finish_reason), output_tokens) - yield sse.message_stop() - yield sse.done() - - def _process_tool_call(self, tc: dict, sse: Any): - """Process a single tool call delta and yield SSE events.""" - tc_index = tc.get("index", 0) - if tc_index < 0: - tc_index = len(sse.blocks.tool_indices) - - fn_delta = tc.get("function", {}) - incoming_name = fn_delta.get("name") - if incoming_name is not None: - sse.blocks.register_tool_name(tc_index, incoming_name) - - if tc_index not in sse.blocks.tool_indices: - name = sse.blocks.tool_names.get(tc_index, "") - if name or tc.get("id"): - tool_id = tc.get("id") or f"tool_{uuid.uuid4()}" - yield sse.start_tool_block(tc_index, tool_id, name) - sse.blocks.tool_started[tc_index] = True - elif not sse.blocks.tool_started.get(tc_index) and sse.blocks.tool_names.get( - tc_index - ): - tool_id = tc.get("id") or f"tool_{uuid.uuid4()}" - name = sse.blocks.tool_names[tc_index] - yield sse.start_tool_block(tc_index, tool_id, name) - sse.blocks.tool_started[tc_index] = True - - args = fn_delta.get("arguments", "") - if args: - if not sse.blocks.tool_started.get(tc_index): - tool_id = tc.get("id") or f"tool_{uuid.uuid4()}" - name = sse.blocks.tool_names.get(tc_index, "tool_call") or "tool_call" - yield sse.start_tool_block(tc_index, tool_id, name) - sse.blocks.tool_started[tc_index] = True - - current_name = sse.blocks.tool_names.get(tc_index, "") - if current_name == "Task": - parsed = sse.blocks.buffer_task_args(tc_index, args) - if parsed is not None: - yield sse.emit_tool_delta(tc_index, json.dumps(parsed)) - return - - yield sse.emit_tool_delta(tc_index, args) - - def _flush_task_arg_buffers(self, sse: Any): - """Emit buffered Task args as a single JSON delta (best-effort).""" - for tool_index, out in sse.blocks.flush_task_arg_buffers(): - yield sse.emit_tool_delta(tool_index, out) diff --git a/providers/lmstudio/request.py b/providers/lmstudio/request.py index aac8cec6..dd87e3a8 100644 --- a/providers/lmstudio/request.py +++ b/providers/lmstudio/request.py @@ -2,7 +2,7 @@ from typing import Any, Dict -from providers.nvidia_nim.utils.message_converter import AnthropicToOpenAIConverter +from providers.common.message_converter import AnthropicToOpenAIConverter from loguru import logger diff --git a/providers/nvidia_nim/client.py b/providers/nvidia_nim/client.py index 3898c70b..4acdf201 100644 --- a/providers/nvidia_nim/client.py +++ b/providers/nvidia_nim/client.py @@ -1,310 +1,27 @@ """NVIDIA NIM provider implementation.""" -import json -import uuid -from typing import Any, AsyncIterator +from typing import Any -import httpx -from loguru import logger -from openai import AsyncOpenAI +from config.nim import NimSettings +from providers.openai_compat import OpenAICompatibleProvider +from providers.base import ProviderConfig -from providers.base import BaseProvider, ProviderConfig -from providers.rate_limit import GlobalRateLimiter from .request import build_request_body -from .errors import map_error -from .utils import ( - SSEBuilder, - map_stop_reason, - ThinkTagParser, - HeuristicToolParser, - ContentType, -) -class NvidiaNimProvider(BaseProvider): +class NvidiaNimProvider(OpenAICompatibleProvider): """NVIDIA NIM provider using official OpenAI client.""" - def __init__(self, config: ProviderConfig): - super().__init__(config) - self._api_key = config.api_key - self._base_url = ( - config.base_url or "https://integrate.api.nvidia.com/v1" - ).rstrip("/") - self._nim_settings = config.nim_settings - self._global_rate_limiter = GlobalRateLimiter.get_instance( - rate_limit=config.rate_limit, - rate_window=config.rate_window, - ) - self._client = AsyncOpenAI( - api_key=self._api_key, - base_url=self._base_url, - max_retries=0, - timeout=httpx.Timeout( - config.http_read_timeout, - connect=config.http_connect_timeout, - read=config.http_read_timeout, - write=config.http_write_timeout, - ), + def __init__(self, config: ProviderConfig, *, nim_settings: NimSettings): + super().__init__( + config, + provider_name="NIM", + base_url=config.base_url or "https://integrate.api.nvidia.com/v1", + api_key=config.api_key, + nim_settings=nim_settings, ) def _build_request_body(self, request: Any) -> dict: """Internal helper for tests and shared building.""" + assert self._nim_settings is not None return build_request_body(request, self._nim_settings) - - async def stream_response( - self, - request: Any, - input_tokens: int = 0, - *, - request_id: str | None = None, - ) -> AsyncIterator[str]: - """Stream response in Anthropic SSE format.""" - with logger.contextualize(request_id=request_id): - async for event in self._stream_response_impl( - request, input_tokens, request_id - ): - yield event - - async def _stream_response_impl( - self, - request: Any, - input_tokens: int, - request_id: str | None, - ) -> AsyncIterator[str]: - """Internal streaming implementation with context bound.""" - message_id = f"msg_{uuid.uuid4()}" - sse = SSEBuilder(message_id, request.model, input_tokens) - - body = self._build_request_body(request) - req_tag = f" request_id={request_id}" if request_id else "" - logger.info( - "NIM_STREAM:%s model=%s msgs=%d tools=%d", - req_tag, - body.get("model"), - len(body.get("messages", [])), - len(body.get("tools", [])), - ) - - yield sse.message_start() - - think_parser = ThinkTagParser() - heuristic_parser = HeuristicToolParser() - - finish_reason = None - usage_info = None - error_occurred = False - error_message = "" - - try: - stream = await self._global_rate_limiter.execute_with_retry( - self._client.chat.completions.create, **body, stream=True - ) - async for chunk in stream: - # OpenAI client returns objects, not JSON - if getattr(chunk, "usage", None): - usage_info = chunk.usage - - if not chunk.choices: - continue - - choice = chunk.choices[0] - delta = choice.delta - if delta is None: - continue - - if choice.finish_reason: - finish_reason = choice.finish_reason - logger.debug(f"NIM finish_reason: {finish_reason}") - - # Handle reasoning content from delta - reasoning = getattr(delta, "reasoning_content", None) - if reasoning: - for event in sse.ensure_thinking_block(): - yield event - yield sse.emit_thinking_delta(reasoning) - - # Handle text content - if delta.content: - for part in think_parser.feed(delta.content): - if part.type == ContentType.THINKING: - for event in sse.ensure_thinking_block(): - yield event - yield sse.emit_thinking_delta(part.content) - else: - filtered_text, detected_tools = heuristic_parser.feed( - part.content - ) - - if filtered_text: - for event in sse.ensure_text_block(): - yield event - yield sse.emit_text_delta(filtered_text) - - for tool_use in detected_tools: - for event in sse.close_content_blocks(): - yield event - - block_idx = sse.blocks.allocate_index() - if tool_use.get("name") == "Task" and isinstance( - tool_use.get("input"), dict - ): - tool_use["input"]["run_in_background"] = False - yield sse.content_block_start( - block_idx, - "tool_use", - id=tool_use["id"], - name=tool_use["name"], - ) - yield sse.content_block_delta( - block_idx, - "input_json_delta", - json.dumps(tool_use["input"]), - ) - yield sse.content_block_stop(block_idx) - - # Handle native tool calls - if delta.tool_calls: - for event in sse.close_content_blocks(): - yield event - for tc in delta.tool_calls: - # Convert OpenAI tool call object to dict for existing logic - tc_info = { - "index": tc.index, - "id": tc.id, - "function": { - "name": tc.function.name, - "arguments": tc.function.arguments, - }, - } - for event in self._process_tool_call(tc_info, sse): - yield event - - except Exception as e: - req_tag = f" request_id={request_id}" if request_id else "" - logger.error("NIM_ERROR:%s %s: %s", req_tag, type(e).__name__, e) - mapped_e = map_error(e) - error_occurred = True - error_message = str(mapped_e) - logger.info( - "NIM_STREAM: Emitting SSE error event for %s%s", - type(e).__name__, - req_tag, - ) - # Ensure open blocks are closed before emitting error to follow Anthropic protocol - for event in sse.close_content_blocks(): - yield event - for event in sse.emit_error(error_message): - yield event - - # Flush remaining content - remaining = think_parser.flush() - if remaining: - if remaining.type == ContentType.THINKING: - for event in sse.ensure_thinking_block(): - yield event - yield sse.emit_thinking_delta(remaining.content) - else: - for event in sse.ensure_text_block(): - yield event - yield sse.emit_text_delta(remaining.content) - - for tool_use in heuristic_parser.flush(): - for event in sse.close_content_blocks(): - yield event - - block_idx = sse.blocks.allocate_index() - yield sse.content_block_start( - block_idx, - "tool_use", - id=tool_use["id"], - name=tool_use["name"], - ) - if tool_use.get("name") == "Task" and isinstance( - tool_use.get("input"), dict - ): - tool_use["input"]["run_in_background"] = False - yield sse.content_block_delta( - block_idx, - "input_json_delta", - json.dumps(tool_use["input"]), - ) - yield sse.content_block_stop(block_idx) - - if ( - not error_occurred - and sse.blocks.text_index == -1 - and not sse.blocks.tool_indices - ): - for event in sse.ensure_text_block(): - yield event - yield sse.emit_text_delta(" ") - - # Flush buffered Task args before closing tool blocks. - for event in self._flush_task_arg_buffers(sse): - yield event - - for event in sse.close_all_blocks(): - yield event - - output_tokens = ( - usage_info.completion_tokens - if usage_info and hasattr(usage_info, "completion_tokens") - else sse.estimate_output_tokens() - ) - if usage_info and hasattr(usage_info, "prompt_tokens"): - provider_input = usage_info.prompt_tokens - if isinstance(provider_input, int): - diff = provider_input - input_tokens - logger.debug( - f"TOKEN_ESTIMATE: our={input_tokens} provider={provider_input} diff={diff:+d}" - ) - yield sse.message_delta(map_stop_reason(finish_reason), output_tokens) - yield sse.message_stop() - yield sse.done() - - def _process_tool_call(self, tc: dict, sse: Any): - """Process a single tool call delta and yield SSE events.""" - tc_index = tc.get("index", 0) - if tc_index < 0: - tc_index = len(sse.blocks.tool_indices) - - fn_delta = tc.get("function", {}) - incoming_name = fn_delta.get("name") - if incoming_name is not None: - sse.blocks.register_tool_name(tc_index, incoming_name) - - if tc_index not in sse.blocks.tool_indices: - name = sse.blocks.tool_names.get(tc_index, "") - if name or tc.get("id"): - tool_id = tc.get("id") or f"tool_{uuid.uuid4()}" - yield sse.start_tool_block(tc_index, tool_id, name) - sse.blocks.tool_started[tc_index] = True - elif not sse.blocks.tool_started.get(tc_index) and sse.blocks.tool_names.get( - tc_index - ): - tool_id = tc.get("id") or f"tool_{uuid.uuid4()}" - name = sse.blocks.tool_names[tc_index] - yield sse.start_tool_block(tc_index, tool_id, name) - sse.blocks.tool_started[tc_index] = True - - args = fn_delta.get("arguments", "") - if args: - if not sse.blocks.tool_started.get(tc_index): - tool_id = tc.get("id") or f"tool_{uuid.uuid4()}" - name = sse.blocks.tool_names.get(tc_index, "tool_call") or "tool_call" - yield sse.start_tool_block(tc_index, tool_id, name) - sse.blocks.tool_started[tc_index] = True - - current_name = sse.blocks.tool_names.get(tc_index, "") - if current_name == "Task": - parsed = sse.blocks.buffer_task_args(tc_index, args) - if parsed is not None: - yield sse.emit_tool_delta(tc_index, json.dumps(parsed)) - return - - yield sse.emit_tool_delta(tc_index, args) - - def _flush_task_arg_buffers(self, sse: Any): - """Emit buffered Task args as a single JSON delta (best-effort).""" - for tool_index, out in sse.blocks.flush_task_arg_buffers(): - yield sse.emit_tool_delta(tool_index, out) diff --git a/providers/nvidia_nim/errors.py b/providers/nvidia_nim/errors.py index d6b1fa2a..4c70b57a 100644 --- a/providers/nvidia_nim/errors.py +++ b/providers/nvidia_nim/errors.py @@ -1,35 +1,5 @@ -"""Error mapping for NVIDIA NIM provider.""" +"""Error mapping for NVIDIA NIM provider (re-exports from providers.common).""" -import openai +from providers.common.error_mapping import map_error -from providers.exceptions import ( - AuthenticationError, - InvalidRequestError, - RateLimitError, - OverloadedError, - APIError, -) -from providers.rate_limit import GlobalRateLimiter - - -def map_error(e: Exception) -> Exception: - """Map OpenAI exception to specific ProviderError.""" - if isinstance(e, openai.AuthenticationError): - return AuthenticationError(str(e), raw_error=str(e)) - if isinstance(e, openai.RateLimitError): - # Trigger global rate limit block - GlobalRateLimiter.get_instance().set_blocked(60) # Default 60s cooldown - return RateLimitError(str(e), raw_error=str(e)) - if isinstance(e, openai.BadRequestError): - return InvalidRequestError(str(e), raw_error=str(e)) - if isinstance(e, openai.InternalServerError): - message = str(e) - if "overloaded" in message.lower() or "capacity" in message.lower(): - return OverloadedError(message, raw_error=str(e)) - return APIError(message, status_code=500, raw_error=str(e)) - if isinstance(e, openai.APIError): - return APIError( - str(e), status_code=getattr(e, "status_code", 500), raw_error=str(e) - ) - - return e +__all__ = ["map_error"] diff --git a/providers/nvidia_nim/request.py b/providers/nvidia_nim/request.py index 140f19dd..dac44a1c 100644 --- a/providers/nvidia_nim/request.py +++ b/providers/nvidia_nim/request.py @@ -3,7 +3,7 @@ from typing import Any, Dict from config.nim import NimSettings -from .utils.message_converter import AnthropicToOpenAIConverter +from providers.common.message_converter import AnthropicToOpenAIConverter from loguru import logger diff --git a/providers/nvidia_nim/utils/__init__.py b/providers/nvidia_nim/utils/__init__.py index 634b6e51..295c90c9 100644 --- a/providers/nvidia_nim/utils/__init__.py +++ b/providers/nvidia_nim/utils/__init__.py @@ -1,13 +1,13 @@ -"""Utility modules for providers.""" +"""Utility modules for providers (re-exports from providers.common).""" -from .sse_builder import SSEBuilder, ContentBlockManager, map_stop_reason -from .think_parser import ( +from providers.common import ( + SSEBuilder, + ContentBlockManager, + map_stop_reason, ThinkTagParser, ContentType, ContentChunk, -) -from .heuristic_tool_parser import HeuristicToolParser -from .message_converter import ( + HeuristicToolParser, AnthropicToOpenAIConverter, get_block_attr, get_block_type, diff --git a/providers/open_router/client.py b/providers/open_router/client.py index 53969ff5..79a126b6 100644 --- a/providers/open_router/client.py +++ b/providers/open_router/client.py @@ -1,23 +1,10 @@ """OpenRouter provider implementation.""" -import json -import uuid -from typing import Any, AsyncIterator +from typing import Any, Iterator -import httpx -from loguru import logger -from openai import AsyncOpenAI - -from providers.base import BaseProvider, ProviderConfig -from providers.rate_limit import GlobalRateLimiter -from providers.nvidia_nim.errors import map_error -from providers.nvidia_nim.utils import ( - SSEBuilder, - map_stop_reason, - ThinkTagParser, - HeuristicToolParser, - ContentType, -) +from providers.openai_compat import OpenAICompatibleProvider +from providers.base import ProviderConfig +from providers.common import SSEBuilder from .request import build_request_body @@ -25,296 +12,28 @@ from .request import build_request_body OPENROUTER_BASE_URL = "https://openrouter.ai/api/v1" -class OpenRouterProvider(BaseProvider): +class OpenRouterProvider(OpenAICompatibleProvider): """OpenRouter provider using OpenAI-compatible API.""" def __init__(self, config: ProviderConfig): - super().__init__(config) - self._api_key = config.api_key - self._base_url = (config.base_url or OPENROUTER_BASE_URL).rstrip("/") - self._global_rate_limiter = GlobalRateLimiter.get_instance( - rate_limit=config.rate_limit, - rate_window=config.rate_window, - ) - self._client = AsyncOpenAI( - api_key=self._api_key, - base_url=self._base_url, - max_retries=0, - timeout=httpx.Timeout( - config.http_read_timeout, - connect=config.http_connect_timeout, - read=config.http_read_timeout, - write=config.http_write_timeout, - ), + super().__init__( + config, + provider_name="OPENROUTER", + base_url=config.base_url or OPENROUTER_BASE_URL, + api_key=config.api_key, ) def _build_request_body(self, request: Any) -> dict: """Internal helper for tests and shared building.""" return build_request_body(request) - async def stream_response( - self, - request: Any, - input_tokens: int = 0, - *, - request_id: str | None = None, - ) -> AsyncIterator[str]: - """Stream response in Anthropic SSE format.""" - with logger.contextualize(request_id=request_id): - async for event in self._stream_response_impl( - request, input_tokens, request_id - ): - yield event - - async def _stream_response_impl( - self, - request: Any, - input_tokens: int, - request_id: str | None, - ) -> AsyncIterator[str]: - """Internal streaming implementation with context bound.""" - message_id = f"msg_{uuid.uuid4()}" - sse = SSEBuilder(message_id, request.model, input_tokens) - - body = self._build_request_body(request) - req_tag = f" request_id={request_id}" if request_id else "" - logger.info( - "OPENROUTER_STREAM:%s model=%s msgs=%d tools=%d", - req_tag, - body.get("model"), - len(body.get("messages", [])), - len(body.get("tools", [])), - ) - - yield sse.message_start() - - think_parser = ThinkTagParser() - heuristic_parser = HeuristicToolParser() - - finish_reason = None - usage_info = None - error_occurred = False - error_message = "" - - try: - stream = await self._global_rate_limiter.execute_with_retry( - self._client.chat.completions.create, **body, stream=True - ) - async for chunk in stream: - if getattr(chunk, "usage", None): - usage_info = chunk.usage - - if not chunk.choices: - continue - - choice = chunk.choices[0] - delta = choice.delta - if delta is None: - continue - - if choice.finish_reason: - finish_reason = choice.finish_reason - logger.debug("OPENROUTER finish_reason: %s", finish_reason) - - # Handle reasoning_content (OpenRouter/OpenAI extended format) - reasoning = getattr(delta, "reasoning_content", None) - if reasoning: + def _handle_extra_reasoning(self, delta: Any, sse: SSEBuilder) -> Iterator[str]: + """Handle reasoning_details for StepFun models.""" + reasoning_details = getattr(delta, "reasoning_details", None) + if reasoning_details and isinstance(reasoning_details, list): + for item in reasoning_details: + text = item.get("text", "") if isinstance(item, dict) else "" + if text: for event in sse.ensure_thinking_block(): yield event - yield sse.emit_thinking_delta(reasoning) - - # Handle reasoning_details (e.g. stepfun models) - reasoning_details = getattr(delta, "reasoning_details", None) - if reasoning_details and isinstance(reasoning_details, list): - for item in reasoning_details: - text = item.get("text", "") if isinstance(item, dict) else "" - if text: - for event in sse.ensure_thinking_block(): - yield event - yield sse.emit_thinking_delta(text) - - # Handle text content - if delta.content: - for part in think_parser.feed(delta.content): - if part.type == ContentType.THINKING: - for event in sse.ensure_thinking_block(): - yield event - yield sse.emit_thinking_delta(part.content) - else: - filtered_text, detected_tools = heuristic_parser.feed( - part.content - ) - - if filtered_text: - for event in sse.ensure_text_block(): - yield event - yield sse.emit_text_delta(filtered_text) - - for tool_use in detected_tools: - for event in sse.close_content_blocks(): - yield event - - block_idx = sse.blocks.allocate_index() - if tool_use.get("name") == "Task" and isinstance( - tool_use.get("input"), dict - ): - tool_use["input"]["run_in_background"] = False - yield sse.content_block_start( - block_idx, - "tool_use", - id=tool_use["id"], - name=tool_use["name"], - ) - yield sse.content_block_delta( - block_idx, - "input_json_delta", - json.dumps(tool_use["input"]), - ) - yield sse.content_block_stop(block_idx) - - # Handle native tool calls - if delta.tool_calls: - for event in sse.close_content_blocks(): - yield event - for tc in delta.tool_calls: - tc_info = { - "index": tc.index, - "id": tc.id, - "function": { - "name": tc.function.name, - "arguments": tc.function.arguments, - }, - } - for event in self._process_tool_call(tc_info, sse): - yield event - - except Exception as e: - req_tag = f" request_id={request_id}" if request_id else "" - logger.error("OPENROUTER_ERROR:%s %s: %s", req_tag, type(e).__name__, e) - mapped_e = map_error(e) - error_occurred = True - error_message = str(mapped_e) - logger.info( - "OPENROUTER_STREAM: Emitting SSE error event for %s%s", - type(e).__name__, - req_tag, - ) - for event in sse.close_content_blocks(): - yield event - for event in sse.emit_error(error_message): - yield event - - # Flush remaining content - remaining = think_parser.flush() - if remaining: - if remaining.type == ContentType.THINKING: - for event in sse.ensure_thinking_block(): - yield event - yield sse.emit_thinking_delta(remaining.content) - else: - for event in sse.ensure_text_block(): - yield event - yield sse.emit_text_delta(remaining.content) - - for tool_use in heuristic_parser.flush(): - for event in sse.close_content_blocks(): - yield event - - block_idx = sse.blocks.allocate_index() - yield sse.content_block_start( - block_idx, - "tool_use", - id=tool_use["id"], - name=tool_use["name"], - ) - if tool_use.get("name") == "Task" and isinstance( - tool_use.get("input"), dict - ): - tool_use["input"]["run_in_background"] = False - yield sse.content_block_delta( - block_idx, - "input_json_delta", - json.dumps(tool_use["input"]), - ) - yield sse.content_block_stop(block_idx) - - if ( - not error_occurred - and sse.blocks.text_index == -1 - and not sse.blocks.tool_indices - ): - for event in sse.ensure_text_block(): - yield event - yield sse.emit_text_delta(" ") - - for event in self._flush_task_arg_buffers(sse): - yield event - - for event in sse.close_all_blocks(): - yield event - - output_tokens = ( - usage_info.completion_tokens - if usage_info and hasattr(usage_info, "completion_tokens") - else sse.estimate_output_tokens() - ) - if usage_info and hasattr(usage_info, "prompt_tokens"): - provider_input = usage_info.prompt_tokens - if isinstance(provider_input, int): - diff = provider_input - input_tokens - logger.debug( - "TOKEN_ESTIMATE: our=%d provider=%d diff=%+d", - input_tokens, - provider_input, - diff, - ) - yield sse.message_delta(map_stop_reason(finish_reason), output_tokens) - yield sse.message_stop() - yield sse.done() - - def _process_tool_call(self, tc: dict, sse: Any): - """Process a single tool call delta and yield SSE events.""" - tc_index = tc.get("index", 0) - if tc_index < 0: - tc_index = len(sse.blocks.tool_indices) - - fn_delta = tc.get("function", {}) - incoming_name = fn_delta.get("name") - if incoming_name is not None: - sse.blocks.register_tool_name(tc_index, incoming_name) - - if tc_index not in sse.blocks.tool_indices: - name = sse.blocks.tool_names.get(tc_index, "") - if name or tc.get("id"): - tool_id = tc.get("id") or f"tool_{uuid.uuid4()}" - yield sse.start_tool_block(tc_index, tool_id, name) - sse.blocks.tool_started[tc_index] = True - elif not sse.blocks.tool_started.get(tc_index) and sse.blocks.tool_names.get( - tc_index - ): - tool_id = tc.get("id") or f"tool_{uuid.uuid4()}" - name = sse.blocks.tool_names[tc_index] - yield sse.start_tool_block(tc_index, tool_id, name) - sse.blocks.tool_started[tc_index] = True - - args = fn_delta.get("arguments", "") - if args: - if not sse.blocks.tool_started.get(tc_index): - tool_id = tc.get("id") or f"tool_{uuid.uuid4()}" - name = sse.blocks.tool_names.get(tc_index, "tool_call") or "tool_call" - yield sse.start_tool_block(tc_index, tool_id, name) - sse.blocks.tool_started[tc_index] = True - - current_name = sse.blocks.tool_names.get(tc_index, "") - if current_name == "Task": - parsed = sse.blocks.buffer_task_args(tc_index, args) - if parsed is not None: - yield sse.emit_tool_delta(tc_index, json.dumps(parsed)) - return - - yield sse.emit_tool_delta(tc_index, args) - - def _flush_task_arg_buffers(self, sse: Any): - """Emit buffered Task args as a single JSON delta (best-effort).""" - for tool_index, out in sse.blocks.flush_task_arg_buffers(): - yield sse.emit_tool_delta(tool_index, out) + yield sse.emit_thinking_delta(text) diff --git a/providers/open_router/request.py b/providers/open_router/request.py index 984b3163..3af8763a 100644 --- a/providers/open_router/request.py +++ b/providers/open_router/request.py @@ -2,7 +2,7 @@ from typing import Any, Dict -from providers.nvidia_nim.utils.message_converter import AnthropicToOpenAIConverter +from providers.common.message_converter import AnthropicToOpenAIConverter from loguru import logger diff --git a/providers/openai_compat.py b/providers/openai_compat.py new file mode 100644 index 00000000..ccc465d6 --- /dev/null +++ b/providers/openai_compat.py @@ -0,0 +1,325 @@ +"""Shared base class for OpenAI-compatible providers (NIM, OpenRouter, LM Studio).""" + +import json +import uuid +from typing import Any, AsyncIterator, Iterator, Optional + +import httpx +from loguru import logger +from openai import AsyncOpenAI + +from providers.base import BaseProvider, ProviderConfig +from providers.rate_limit import GlobalRateLimiter +from providers.common import ( + SSEBuilder, + map_stop_reason, + ThinkTagParser, + HeuristicToolParser, + ContentType, + map_error, +) + + +class OpenAICompatibleProvider(BaseProvider): + """Base class for providers using OpenAI-compatible chat completions API.""" + + def __init__( + self, + config: ProviderConfig, + *, + provider_name: str, + base_url: str, + api_key: str, + nim_settings: Optional[Any] = None, + ): + super().__init__(config) + self._provider_name = provider_name + self._api_key = api_key + self._base_url = base_url.rstrip("/") + self._nim_settings = nim_settings + self._global_rate_limiter = GlobalRateLimiter.get_instance( + rate_limit=config.rate_limit, + rate_window=config.rate_window, + ) + self._client = AsyncOpenAI( + api_key=self._api_key, + base_url=self._base_url, + max_retries=0, + timeout=httpx.Timeout( + config.http_read_timeout, + connect=config.http_connect_timeout, + read=config.http_read_timeout, + write=config.http_write_timeout, + ), + ) + + def _build_request_body(self, request: Any) -> dict: + """Build request body. Override in subclasses.""" + raise NotImplementedError + + def _handle_extra_reasoning(self, delta: Any, sse: SSEBuilder) -> Iterator[str]: + """Hook for provider-specific reasoning (e.g. OpenRouter reasoning_details).""" + return iter(()) + + def _process_tool_call(self, tc: dict, sse: Any) -> Iterator[str]: + """Process a single tool call delta and yield SSE events.""" + tc_index = tc.get("index", 0) + if tc_index < 0: + tc_index = len(sse.blocks.tool_indices) + + fn_delta = tc.get("function", {}) + incoming_name = fn_delta.get("name") + if incoming_name is not None: + sse.blocks.register_tool_name(tc_index, incoming_name) + + if tc_index not in sse.blocks.tool_indices: + name = sse.blocks.tool_names.get(tc_index, "") + if name or tc.get("id"): + tool_id = tc.get("id") or f"tool_{uuid.uuid4()}" + yield sse.start_tool_block(tc_index, tool_id, name) + sse.blocks.tool_started[tc_index] = True + elif not sse.blocks.tool_started.get(tc_index) and sse.blocks.tool_names.get( + tc_index + ): + tool_id = tc.get("id") or f"tool_{uuid.uuid4()}" + name = sse.blocks.tool_names[tc_index] + yield sse.start_tool_block(tc_index, tool_id, name) + sse.blocks.tool_started[tc_index] = True + + args = fn_delta.get("arguments", "") + if args: + if not sse.blocks.tool_started.get(tc_index): + tool_id = tc.get("id") or f"tool_{uuid.uuid4()}" + name = sse.blocks.tool_names.get(tc_index, "tool_call") or "tool_call" + yield sse.start_tool_block(tc_index, tool_id, name) + sse.blocks.tool_started[tc_index] = True + + current_name = sse.blocks.tool_names.get(tc_index, "") + if current_name == "Task": + parsed = sse.blocks.buffer_task_args(tc_index, args) + if parsed is not None: + yield sse.emit_tool_delta(tc_index, json.dumps(parsed)) + return + + yield sse.emit_tool_delta(tc_index, args) + + def _flush_task_arg_buffers(self, sse: Any) -> Iterator[str]: + """Emit buffered Task args as a single JSON delta (best-effort).""" + for tool_index, out in sse.blocks.flush_task_arg_buffers(): + yield sse.emit_tool_delta(tool_index, out) + + async def stream_response( + self, + request: Any, + input_tokens: int = 0, + *, + request_id: str | None = None, + ) -> AsyncIterator[str]: + """Stream response in Anthropic SSE format.""" + with logger.contextualize(request_id=request_id): + async for event in self._stream_response_impl( + request, input_tokens, request_id + ): + yield event + + async def _stream_response_impl( + self, + request: Any, + input_tokens: int, + request_id: str | None, + ) -> AsyncIterator[str]: + """Shared streaming implementation.""" + tag = self._provider_name + message_id = f"msg_{uuid.uuid4()}" + sse = SSEBuilder(message_id, request.model, input_tokens) + + body = self._build_request_body(request) + req_tag = f" request_id={request_id}" if request_id else "" + logger.info( + "%s_STREAM:%s model=%s msgs=%d tools=%d", + tag, + req_tag, + body.get("model"), + len(body.get("messages", [])), + len(body.get("tools", [])), + ) + + yield sse.message_start() + + think_parser = ThinkTagParser() + heuristic_parser = HeuristicToolParser() + + finish_reason = None + usage_info = None + error_occurred = False + error_message = "" + + try: + stream = await self._global_rate_limiter.execute_with_retry( + self._client.chat.completions.create, **body, stream=True + ) + async for chunk in stream: + if getattr(chunk, "usage", None): + usage_info = chunk.usage + + if not chunk.choices: + continue + + choice = chunk.choices[0] + delta = choice.delta + if delta is None: + continue + + if choice.finish_reason: + finish_reason = choice.finish_reason + logger.debug("%s finish_reason: %s", tag, finish_reason) + + # Handle reasoning_content (OpenAI extended format) + reasoning = getattr(delta, "reasoning_content", None) + if reasoning: + for event in sse.ensure_thinking_block(): + yield event + yield sse.emit_thinking_delta(reasoning) + + # Provider-specific extra reasoning (e.g. OpenRouter reasoning_details) + for event in self._handle_extra_reasoning(delta, sse): + yield event + + # Handle text content + if delta.content: + for part in think_parser.feed(delta.content): + if part.type == ContentType.THINKING: + for event in sse.ensure_thinking_block(): + yield event + yield sse.emit_thinking_delta(part.content) + else: + filtered_text, detected_tools = heuristic_parser.feed( + part.content + ) + + if filtered_text: + for event in sse.ensure_text_block(): + yield event + yield sse.emit_text_delta(filtered_text) + + for tool_use in detected_tools: + for event in sse.close_content_blocks(): + yield event + + block_idx = sse.blocks.allocate_index() + if tool_use.get("name") == "Task" and isinstance( + tool_use.get("input"), dict + ): + tool_use["input"]["run_in_background"] = False + yield sse.content_block_start( + block_idx, + "tool_use", + id=tool_use["id"], + name=tool_use["name"], + ) + yield sse.content_block_delta( + block_idx, + "input_json_delta", + json.dumps(tool_use["input"]), + ) + yield sse.content_block_stop(block_idx) + + # Handle native tool calls + if delta.tool_calls: + for event in sse.close_content_blocks(): + yield event + for tc in delta.tool_calls: + tc_info = { + "index": tc.index, + "id": tc.id, + "function": { + "name": tc.function.name, + "arguments": tc.function.arguments, + }, + } + for event in self._process_tool_call(tc_info, sse): + yield event + + except Exception as e: + req_tag = f" request_id={request_id}" if request_id else "" + logger.error("%s_ERROR:%s %s: %s", tag, req_tag, type(e).__name__, e) + mapped_e = map_error(e) + error_occurred = True + error_message = str(mapped_e) + logger.info( + "%s_STREAM: Emitting SSE error event for %s%s", + tag, + type(e).__name__, + req_tag, + ) + for event in sse.close_content_blocks(): + yield event + for event in sse.emit_error(error_message): + yield event + + # Flush remaining content + remaining = think_parser.flush() + if remaining: + if remaining.type == ContentType.THINKING: + for event in sse.ensure_thinking_block(): + yield event + yield sse.emit_thinking_delta(remaining.content) + else: + for event in sse.ensure_text_block(): + yield event + yield sse.emit_text_delta(remaining.content) + + for tool_use in heuristic_parser.flush(): + for event in sse.close_content_blocks(): + yield event + + block_idx = sse.blocks.allocate_index() + yield sse.content_block_start( + block_idx, + "tool_use", + id=tool_use["id"], + name=tool_use["name"], + ) + if tool_use.get("name") == "Task" and isinstance( + tool_use.get("input"), dict + ): + tool_use["input"]["run_in_background"] = False + yield sse.content_block_delta( + block_idx, + "input_json_delta", + json.dumps(tool_use["input"]), + ) + yield sse.content_block_stop(block_idx) + + if ( + not error_occurred + and sse.blocks.text_index == -1 + and not sse.blocks.tool_indices + ): + for event in sse.ensure_text_block(): + yield event + yield sse.emit_text_delta(" ") + + for event in self._flush_task_arg_buffers(sse): + yield event + + for event in sse.close_all_blocks(): + yield event + + output_tokens = ( + usage_info.completion_tokens + if usage_info and hasattr(usage_info, "completion_tokens") + else sse.estimate_output_tokens() + ) + if usage_info and hasattr(usage_info, "prompt_tokens"): + provider_input = usage_info.prompt_tokens + if isinstance(provider_input, int): + logger.debug( + "TOKEN_ESTIMATE: our=%d provider=%d diff=%+d", + input_tokens, + provider_input, + provider_input - input_tokens, + ) + yield sse.message_delta(map_stop_reason(finish_reason), output_tokens) + yield sse.message_stop() + yield sse.done() diff --git a/tests/conftest.py b/tests/conftest.py index 6f672116..71f8086c 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -29,13 +29,12 @@ def provider_config(): base_url="https://test.api.nvidia.com/v1", rate_limit=10, rate_window=60, - nim_settings=NimSettings(), ) @pytest.fixture def nim_provider(provider_config): - return NvidiaNimProvider(provider_config) + return NvidiaNimProvider(provider_config, nim_settings=NimSettings()) @pytest.fixture @@ -54,7 +53,6 @@ def lmstudio_provider(provider_config): base_url="http://localhost:1234/v1", rate_limit=provider_config.rate_limit, rate_window=provider_config.rate_window, - nim_settings=provider_config.nim_settings, ) return LMStudioProvider(lmstudio_config) diff --git a/tests/test_converter.py b/tests/test_converter.py index a6dec58c..424dea31 100644 --- a/tests/test_converter.py +++ b/tests/test_converter.py @@ -1,6 +1,6 @@ import json import pytest -from providers.nvidia_nim.utils.message_converter import AnthropicToOpenAIConverter +from providers.common.message_converter import AnthropicToOpenAIConverter # --- Mock Classes --- @@ -273,7 +273,7 @@ def test_convert_mixed_blocks_and_types_and_roles(): def test_get_block_attr_defaults(): # Test helper directly - from providers.nvidia_nim.utils.message_converter import get_block_attr + from providers.common.message_converter import get_block_attr assert get_block_attr({}, "missing", "default") == "default" assert get_block_attr(object(), "missing", "default") == "default" diff --git a/tests/test_dependencies.py b/tests/test_dependencies.py index eeef8155..2aafdb7e 100644 --- a/tests/test_dependencies.py +++ b/tests/test_dependencies.py @@ -127,7 +127,7 @@ async def test_get_provider_passes_http_timeouts_from_settings(): """Provider receives http timeouts from settings when creating client.""" with ( patch("api.dependencies.get_settings") as mock_settings, - patch("providers.nvidia_nim.client.AsyncOpenAI") as mock_openai, + patch("providers.openai_compat.AsyncOpenAI") as mock_openai, ): mock_settings.return_value = _make_mock_settings( http_read_timeout=600.0, diff --git a/tests/test_error_mapping.py b/tests/test_error_mapping.py index c57b1198..8b5ab318 100644 --- a/tests/test_error_mapping.py +++ b/tests/test_error_mapping.py @@ -6,7 +6,7 @@ from unittest.mock import patch, MagicMock import openai from httpx import Response, Request -from providers.nvidia_nim.errors import map_error +from providers.common import map_error from providers.exceptions import ( AuthenticationError, InvalidRequestError, @@ -41,7 +41,7 @@ class TestMapError: def test_rate_limit_error(self): """openai.RateLimitError -> RateLimitError and triggers global block.""" exc = _make_openai_error(openai.RateLimitError, status_code=429) - with patch("providers.nvidia_nim.errors.GlobalRateLimiter") as mock_rl: + with patch("providers.common.error_mapping.GlobalRateLimiter") as mock_rl: mock_instance = MagicMock() mock_rl.get_instance.return_value = mock_instance result = map_error(exc) @@ -117,6 +117,6 @@ class TestMapError: openai.BadRequestError: 400, } exc = _make_openai_error(exc_cls, status_code=status_map[exc_cls]) - with patch("providers.nvidia_nim.errors.GlobalRateLimiter"): + with patch("providers.common.error_mapping.GlobalRateLimiter"): result = map_error(exc) assert isinstance(result, expected_cls) diff --git a/tests/test_lmstudio.py b/tests/test_lmstudio.py index 8b56a83c..2db20299 100644 --- a/tests/test_lmstudio.py +++ b/tests/test_lmstudio.py @@ -7,7 +7,6 @@ from unittest.mock import AsyncMock, MagicMock, patch from providers.base import ProviderConfig from providers.lmstudio import LMStudioProvider from providers.lmstudio.request import LMSTUDIO_DEFAULT_MAX_TOKENS -from config.nim import NimSettings class AsyncStreamMock: @@ -102,14 +101,13 @@ def lmstudio_config(): base_url="http://localhost:1234/v1", rate_limit=10, rate_window=60, - nim_settings=NimSettings(), ) @pytest.fixture(autouse=True) def mock_rate_limiter(): """Mock the global rate limiter to prevent waiting.""" - with patch("providers.lmstudio.client.GlobalRateLimiter") as mock: + with patch("providers.openai_compat.GlobalRateLimiter") as mock: instance = mock.get_instance.return_value instance.wait_if_blocked = AsyncMock(return_value=False) @@ -127,7 +125,7 @@ def lmstudio_provider(lmstudio_config): def test_init(lmstudio_config): """Test provider initialization.""" - with patch("providers.lmstudio.client.AsyncOpenAI") as mock_openai: + with patch("providers.openai_compat.AsyncOpenAI") as mock_openai: provider = LMStudioProvider(lmstudio_config) assert provider._api_key == "lm-studio" assert provider._base_url == "http://localhost:1234/v1" @@ -141,9 +139,8 @@ def test_init_with_empty_api_key(): base_url="http://localhost:1234/v1", rate_limit=10, rate_window=60, - nim_settings=NimSettings(), ) - with patch("providers.lmstudio.client.AsyncOpenAI"): + with patch("providers.openai_compat.AsyncOpenAI"): provider = LMStudioProvider(config) assert provider._api_key == "lm-studio" @@ -157,7 +154,7 @@ def test_init_uses_configurable_timeouts(): http_write_timeout=15.0, http_connect_timeout=5.0, ) - with patch("providers.lmstudio.client.AsyncOpenAI") as mock_openai: + with patch("providers.openai_compat.AsyncOpenAI") as mock_openai: LMStudioProvider(config) call_kwargs = mock_openai.call_args[1] timeout = call_kwargs["timeout"] @@ -471,7 +468,7 @@ class TestLMStudioProcessToolCall: def test_tool_call_with_id(self, lmstudio_provider): """Tool call with id starts a tool block.""" - from providers.nvidia_nim.utils import SSEBuilder + from providers.common import SSEBuilder sse = SSEBuilder("msg_test", "test-model") tc = { @@ -487,7 +484,7 @@ class TestLMStudioProcessToolCall: def test_tool_call_without_id_generates_uuid(self, lmstudio_provider): """Tool call without id generates a uuid-based id.""" - from providers.nvidia_nim.utils import SSEBuilder + from providers.common import SSEBuilder sse = SSEBuilder("msg_test", "test-model") tc = { @@ -501,7 +498,7 @@ class TestLMStudioProcessToolCall: def test_task_tool_forces_background_false(self, lmstudio_provider): """Task tool with run_in_background=true is forced to false.""" - from providers.nvidia_nim.utils import SSEBuilder + from providers.common import SSEBuilder sse = SSEBuilder("msg_test", "test-model") args = json.dumps({"run_in_background": True, "prompt": "test"}) @@ -516,7 +513,7 @@ class TestLMStudioProcessToolCall: def test_task_tool_chunked_args_forces_background_false(self, lmstudio_provider): """Chunked Task args are buffered until valid JSON, then forced to false.""" - from providers.nvidia_nim.utils import SSEBuilder + from providers.common import SSEBuilder sse = SSEBuilder("msg_test", "test-model") tc1 = { @@ -542,7 +539,7 @@ class TestLMStudioProcessToolCall: self, lmstudio_provider, caplog ): """Invalid JSON args for Task tool emits {} on flush and logs a warning.""" - from providers.nvidia_nim.utils import SSEBuilder + from providers.common import SSEBuilder sse = SSEBuilder("msg_test", "test-model") tc = { @@ -561,7 +558,7 @@ class TestLMStudioProcessToolCall: def test_negative_tool_index_fallback(self, lmstudio_provider): """tc_index < 0 uses len(tool_indices) as fallback.""" - from providers.nvidia_nim.utils import SSEBuilder + from providers.common import SSEBuilder sse = SSEBuilder("msg_test", "test-model") tc = { @@ -574,7 +571,7 @@ class TestLMStudioProcessToolCall: def test_tool_args_emitted_as_delta(self, lmstudio_provider): """Arguments are emitted as input_json_delta events.""" - from providers.nvidia_nim.utils import SSEBuilder + from providers.common import SSEBuilder sse = SSEBuilder("msg_test", "test-model") tc = { @@ -588,7 +585,7 @@ class TestLMStudioProcessToolCall: def test_stream_malformed_tool_args_chunked(self, lmstudio_provider): """Chunked tool args that never form valid JSON are flushed with {}.""" - from providers.nvidia_nim.utils import SSEBuilder + from providers.common import SSEBuilder sse = SSEBuilder("msg_test", "test-model") tc1 = { @@ -655,8 +652,7 @@ def test_init_base_url_strips_trailing_slash(): base_url="http://localhost:1234/v1/", rate_limit=10, rate_window=60, - nim_settings=NimSettings(), ) - with patch("providers.lmstudio.client.AsyncOpenAI"): + with patch("providers.openai_compat.AsyncOpenAI"): provider = LMStudioProvider(config) assert provider._base_url == "http://localhost:1234/v1" diff --git a/tests/test_messaging.py b/tests/test_messaging.py index 13782321..cd50eba6 100644 --- a/tests/test_messaging.py +++ b/tests/test_messaging.py @@ -62,7 +62,7 @@ class TestSessionStore: from messaging.session import SessionStore store = SessionStore(storage_path=str(tmp_path / "sessions.json")) - assert store._sessions == {} + assert store._trees == {} # --- Tree Tests --- @@ -127,22 +127,15 @@ class TestSessionStore: # --- Persistence & Edge Cases --- - def test_load_existing_legacy_format(self, tmp_path): - """Test loading legacy session format (int IDs) - backward compat.""" + def test_load_existing_file_with_trees(self, tmp_path): + """Test loading file with trees (legacy sessions ignored).""" from messaging.session import SessionStore data = { - "sessions": { - "s1": { - "session_id": "s1", - "chat_id": 123, # Legacy int - "initial_msg_id": 100, # Legacy int - "last_msg_id": 101, # Legacy int - "created_at": "2024-01-01", - "updated_at": "2024-01-01", - # platform missing -> should default to telegram - } - } + "sessions": {}, + "trees": {"r1": {"root_id": "r1", "nodes": {"r1": {}}}}, + "node_to_tree": {"r1": "r1"}, + "message_log": {}, } p = tmp_path / "sessions.json" @@ -150,11 +143,7 @@ class TestSessionStore: json.dump(data, f) store = SessionStore(storage_path=str(p)) - # Legacy sessions are loaded for backward compat; verify conversion - assert "s1" in store._sessions - rec = store._sessions["s1"] - assert rec.chat_id == "123" # Converted to str - assert rec.platform == "telegram" # Defaulted + assert store.get_tree("r1") is not None def test_load_corrupt_file(self, tmp_path): """Test loading corrupt/invalid json file.""" @@ -166,7 +155,7 @@ class TestSessionStore: # Should log error and start empty, avoiding crash store = SessionStore(storage_path=str(p)) - assert store._sessions == {} + assert store._trees == {} def test_save_error_handling(self, tmp_path): """Test error during save.""" diff --git a/tests/test_nvidia_nim.py b/tests/test_nvidia_nim.py index e1c4e2da..253b388f 100644 --- a/tests/test_nvidia_nim.py +++ b/tests/test_nvidia_nim.py @@ -38,7 +38,7 @@ class MockRequest: @pytest.fixture(autouse=True) def mock_rate_limiter(): """Mock the global rate limiter to prevent waiting.""" - with patch("providers.nvidia_nim.client.GlobalRateLimiter") as mock: + with patch("providers.openai_compat.GlobalRateLimiter") as mock: instance = mock.get_instance.return_value instance.wait_if_blocked = AsyncMock(return_value=False) @@ -53,8 +53,10 @@ def mock_rate_limiter(): @pytest.mark.asyncio async def test_init(provider_config): """Test provider initialization.""" - with patch("providers.nvidia_nim.client.AsyncOpenAI") as mock_openai: - provider = NvidiaNimProvider(provider_config) + with patch("providers.openai_compat.AsyncOpenAI") as mock_openai: + from config.nim import NimSettings + + provider = NvidiaNimProvider(provider_config, nim_settings=NimSettings()) assert provider._api_key == "test_key" assert provider._base_url == "https://test.api.nvidia.com/v1" mock_openai.assert_called_once() @@ -64,6 +66,7 @@ async def test_init(provider_config): async def test_init_uses_configurable_timeouts(): """Test that provider passes configurable read/write/connect timeouts to client.""" from providers.base import ProviderConfig + from config.nim import NimSettings config = ProviderConfig( api_key="test_key", @@ -72,8 +75,8 @@ async def test_init_uses_configurable_timeouts(): http_write_timeout=15.0, http_connect_timeout=5.0, ) - with patch("providers.nvidia_nim.client.AsyncOpenAI") as mock_openai: - NvidiaNimProvider(config) + with patch("providers.openai_compat.AsyncOpenAI") as mock_openai: + NvidiaNimProvider(config, nim_settings=NimSettings()) call_kwargs = mock_openai.call_args[1] timeout = call_kwargs["timeout"] assert timeout.read == 600.0 diff --git a/tests/test_open_router.py b/tests/test_open_router.py index 7f3b58ce..4fcd3604 100644 --- a/tests/test_open_router.py +++ b/tests/test_open_router.py @@ -6,7 +6,6 @@ from unittest.mock import MagicMock, AsyncMock, patch from providers.open_router import OpenRouterProvider from providers.open_router.request import OPENROUTER_DEFAULT_MAX_TOKENS from providers.base import ProviderConfig -from config.nim import NimSettings class MockMessage: @@ -39,14 +38,13 @@ def open_router_config(): base_url="https://openrouter.ai/api/v1", rate_limit=10, rate_window=60, - nim_settings=NimSettings(), ) @pytest.fixture(autouse=True) def mock_rate_limiter(): """Mock the global rate limiter to prevent waiting.""" - with patch("providers.open_router.client.GlobalRateLimiter") as mock: + with patch("providers.openai_compat.GlobalRateLimiter") as mock: instance = mock.get_instance.return_value instance.wait_if_blocked = AsyncMock(return_value=False) @@ -64,7 +62,7 @@ def open_router_provider(open_router_config): def test_init(open_router_config): """Test provider initialization.""" - with patch("providers.open_router.client.AsyncOpenAI") as mock_openai: + with patch("providers.openai_compat.AsyncOpenAI") as mock_openai: provider = OpenRouterProvider(open_router_config) assert provider._api_key == "test_openrouter_key" assert provider._base_url == "https://openrouter.ai/api/v1" @@ -80,7 +78,7 @@ def test_init_uses_configurable_timeouts(): http_write_timeout=15.0, http_connect_timeout=5.0, ) - with patch("providers.open_router.client.AsyncOpenAI") as mock_openai: + with patch("providers.openai_compat.AsyncOpenAI") as mock_openai: OpenRouterProvider(config) call_kwargs = mock_openai.call_args[1] timeout = call_kwargs["timeout"] diff --git a/tests/test_parsers.py b/tests/test_parsers.py index 89e3c537..4a2ea19d 100644 --- a/tests/test_parsers.py +++ b/tests/test_parsers.py @@ -1,7 +1,6 @@ import pytest -from providers.nvidia_nim.utils.think_parser import ThinkTagParser, ContentType -from providers.nvidia_nim.utils.heuristic_tool_parser import HeuristicToolParser +from providers.common import ThinkTagParser, ContentType, HeuristicToolParser def test_think_tag_parser_basic(): diff --git a/tests/test_session_store_edge_cases.py b/tests/test_session_store_edge_cases.py index ab013324..05962d3b 100644 --- a/tests/test_session_store_edge_cases.py +++ b/tests/test_session_store_edge_cases.py @@ -24,7 +24,6 @@ class TestSessionStoreLoadEdgeCases: f.write("{invalid json") store = SessionStore(storage_path=path) - assert len(store._sessions) == 0 assert len(store._trees) == 0 def test_load_truncated_json(self, tmp_path): @@ -34,7 +33,7 @@ class TestSessionStoreLoadEdgeCases: f.write('{"sessions": {"s1": {"session_id": "s1"') store = SessionStore(storage_path=path) - assert len(store._sessions) == 0 + assert len(store._trees) == 0 def test_load_empty_file(self, tmp_path): """Empty file is handled gracefully.""" @@ -43,16 +42,16 @@ class TestSessionStoreLoadEdgeCases: f.write("") store = SessionStore(storage_path=path) - assert len(store._sessions) == 0 + assert len(store._trees) == 0 def test_load_nonexistent_file(self, tmp_path): """Non-existent file starts with empty state.""" path = str(tmp_path / "nonexistent.json") store = SessionStore(storage_path=path) - assert len(store._sessions) == 0 + assert len(store._trees) == 0 - def test_load_legacy_int_fields(self, tmp_path): - """Legacy format with int chat_id/msg_id is converted to string.""" + def test_load_legacy_sessions_ignored(self, tmp_path): + """Legacy sessions in file are ignored; trees and message_log load.""" path = str(tmp_path / "sessions.json") data = { "sessions": { @@ -66,17 +65,15 @@ class TestSessionStoreLoadEdgeCases: "updated_at": "2025-01-01T00:00:00+00:00", } }, - "trees": {}, - "node_to_tree": {}, + "trees": {"r1": {"root_id": "r1", "nodes": {"r1": {}}}}, + "node_to_tree": {"r1": "r1"}, + "message_log": {}, } with open(path, "w") as f: json.dump(data, f) store = SessionStore(storage_path=path) - record = store._sessions["s1"] - assert record.chat_id == "12345" - assert record.initial_msg_id == "100" - assert record.last_msg_id == "200" + assert store.get_tree("r1") is not None class TestSessionStoreSaveEdgeCases: @@ -131,13 +128,11 @@ class TestSessionStoreClearAll: with open(path, "r", encoding="utf-8") as f: data = json.load(f) - assert data["sessions"] == {} assert data["trees"] == {} assert data["node_to_tree"] == {} assert data["message_log"] == {} store2 = SessionStore(storage_path=path) - assert len(store2._sessions) == 0 assert len(store2._trees) == 0 def test_message_log_persists_and_dedups(self, tmp_path): diff --git a/tests/test_sse_builder.py b/tests/test_sse_builder.py index 1ed58f92..0e8ca1e3 100644 --- a/tests/test_sse_builder.py +++ b/tests/test_sse_builder.py @@ -4,7 +4,7 @@ import json import pytest from unittest.mock import patch -from providers.nvidia_nim.utils.sse_builder import ( +from providers.common.sse_builder import ( SSEBuilder, ContentBlockManager, map_stop_reason, @@ -369,7 +369,7 @@ class TestSSEBuilderTokenEstimation: builder.start_text_block() builder.emit_text_delta("a" * 100) # 100 chars -> ~25 tokens - with patch("providers.nvidia_nim.utils.sse_builder.ENCODER", None): + with patch("providers.common.sse_builder.ENCODER", None): tokens = builder.estimate_output_tokens() assert tokens == 25 # 100 // 4 @@ -379,7 +379,7 @@ class TestSSEBuilderTokenEstimation: builder.start_tool_block(0, "t1", "Read") builder.emit_tool_delta(0, '{"path":"test.py"}') - with patch("providers.nvidia_nim.utils.sse_builder.ENCODER", None): + with patch("providers.common.sse_builder.ENCODER", None): tokens = builder.estimate_output_tokens() # 1 tool * 50 = 50 assert tokens == 50 diff --git a/tests/test_streaming_errors.py b/tests/test_streaming_errors.py index 50bb017e..1b5c0fda 100644 --- a/tests/test_streaming_errors.py +++ b/tests/test_streaming_errors.py @@ -33,9 +33,8 @@ def _make_provider(): base_url="https://test.api.nvidia.com/v1", rate_limit=10, rate_window=60, - nim_settings=NimSettings(), ) - return NvidiaNimProvider(config) + return NvidiaNimProvider(config, nim_settings=NimSettings()) def _make_request(model="test-model", stream=True): @@ -270,7 +269,7 @@ class TestProcessToolCall: def test_tool_call_with_id(self): """Tool call with id starts a tool block.""" provider = _make_provider() - from providers.nvidia_nim.utils import SSEBuilder + from providers.common import SSEBuilder sse = SSEBuilder("msg_test", "test-model") tc = { @@ -287,7 +286,7 @@ class TestProcessToolCall: def test_tool_call_without_id_generates_uuid(self): """Tool call without id generates a uuid-based id.""" provider = _make_provider() - from providers.nvidia_nim.utils import SSEBuilder + from providers.common import SSEBuilder sse = SSEBuilder("msg_test", "test-model") tc = { @@ -302,7 +301,7 @@ class TestProcessToolCall: def test_task_tool_forces_background_false(self): """Task tool with run_in_background=true is forced to false.""" provider = _make_provider() - from providers.nvidia_nim.utils import SSEBuilder + from providers.common import SSEBuilder sse = SSEBuilder("msg_test", "test-model") args = json.dumps({"run_in_background": True, "prompt": "test"}) @@ -319,7 +318,7 @@ class TestProcessToolCall: def test_task_tool_chunked_args_forces_background_false(self): """Chunked Task args are buffered until valid JSON, then forced to false.""" provider = _make_provider() - from providers.nvidia_nim.utils import SSEBuilder + from providers.common import SSEBuilder sse = SSEBuilder("msg_test", "test-model") tc1 = { @@ -344,7 +343,7 @@ class TestProcessToolCall: def test_task_tool_invalid_json_logs_warning_on_flush(self, caplog): """Invalid JSON args for Task tool emits {} on flush and logs a warning.""" provider = _make_provider() - from providers.nvidia_nim.utils import SSEBuilder + from providers.common import SSEBuilder sse = SSEBuilder("msg_test", "test-model") tc = { @@ -364,7 +363,7 @@ class TestProcessToolCall: def test_negative_tool_index_fallback(self): """tc_index < 0 uses len(tool_indices) as fallback.""" provider = _make_provider() - from providers.nvidia_nim.utils import SSEBuilder + from providers.common import SSEBuilder sse = SSEBuilder("msg_test", "test-model") tc = { @@ -379,7 +378,7 @@ class TestProcessToolCall: def test_tool_args_emitted_as_delta(self): """Arguments are emitted as input_json_delta events.""" provider = _make_provider() - from providers.nvidia_nim.utils import SSEBuilder + from providers.common import SSEBuilder sse = SSEBuilder("msg_test", "test-model") tc = { @@ -496,7 +495,7 @@ class TestStreamChunkEdgeCases: def test_stream_malformed_tool_args_chunked(self): """Chunked tool args that never form valid JSON are flushed with {}.""" provider = _make_provider() - from providers.nvidia_nim.utils import SSEBuilder + from providers.common import SSEBuilder sse = SSEBuilder("msg_test", "test-model") tc1 = { diff --git a/tests/test_subagent_interception.py b/tests/test_subagent_interception.py index d2fef13b..b3fc3b2b 100644 --- a/tests/test_subagent_interception.py +++ b/tests/test_subagent_interception.py @@ -2,15 +2,16 @@ import json import pytest from unittest.mock import MagicMock from providers.nvidia_nim import NvidiaNimProvider -from providers.nvidia_nim.utils.sse_builder import ContentBlockManager +from providers.common import ContentBlockManager from providers.base import ProviderConfig +from config.nim import NimSettings @pytest.mark.asyncio async def test_task_tool_interception(): # Setup provider config = ProviderConfig(api_key="test") - provider = NvidiaNimProvider(config) + provider = NvidiaNimProvider(config, nim_settings=NimSettings()) # Mock request and sse builder with real ContentBlockManager request = MagicMock()