Merge pull request #32 from Alishahryar1/claude/code-review-architecture-tRi3B

Add architecture improvement plan from code review
This commit is contained in:
Ali Khokhar
2026-02-16 18:07:04 -08:00
committed by GitHub
39 changed files with 1075 additions and 1156 deletions
+31 -21
View File
@@ -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.
- Prefer built-in tools (grep, read_file, etc.) over manual workflows. Check tool availability before use.
+31 -21
View File
@@ -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.
- Prefer built-in tools (grep, read_file, etc.) over manual workflows. Check tool availability before use.
+450
View File
@@ -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** |
+1 -4
View File
@@ -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,
+1 -1
View File
@@ -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,
}
+5 -5
View File
@@ -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:
+14 -64
View File
@@ -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()
+20 -10
View File
@@ -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)
+12 -4
View File
@@ -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]:
+4 -4
View File
@@ -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
+3 -6
View File
@@ -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
+25
View File
@@ -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",
]
+35
View File
@@ -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
@@ -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)
+9 -291
View File
@@ -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)
+1 -1
View File
@@ -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
+13 -296
View File
@@ -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)
+3 -33
View File
@@ -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"]
+1 -1
View File
@@ -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
+6 -6
View File
@@ -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,
+18 -299
View File
@@ -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)
+1 -1
View File
@@ -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
+325
View File
@@ -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()
+1 -3
View File
@@ -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)
+2 -2
View File
@@ -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"
+1 -1
View File
@@ -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,
+3 -3
View File
@@ -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)
+13 -17
View File
@@ -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"
+9 -20
View File
@@ -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."""
+8 -5
View File
@@ -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
+3 -5
View File
@@ -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"]
+1 -2
View File
@@ -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():
+9 -14
View File
@@ -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):
+3 -3
View File
@@ -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
+9 -10
View File
@@ -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 = {
+3 -2
View File
@@ -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()