mirror of
https://github.com/Alishahryar1/free-claude-code.git
synced 2026-07-03 14:05:26 +02:00
Merge pull request #32 from Alishahryar1/claude/code-review-architecture-tRi3B
Add architecture improvement plan from code review
This commit is contained in:
@@ -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.
|
||||
@@ -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.
|
||||
@@ -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
@@ -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
@@ -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,
|
||||
}
|
||||
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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]:
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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
|
||||
+1
-1
@@ -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)
|
||||
@@ -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)
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
@@ -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"]
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
@@ -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
@@ -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."""
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,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():
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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 = {
|
||||
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user