refactor(task): 拆分 TaskService 并重构任务生命周期

- 将 task 公共协议/上下文/异常/schema 下沉到 core,并迁移 polling 目录
- 新增 execute/submit/video 子模块,拆分同步执行、异步提交流程、错误处理与视频任务操作
- 收敛 TaskService 为门面编排,内部委派到 SyncTaskExecutionService、AsyncTaskSubmitService、VideoTaskOperationsService
- 重构 main 生命周期管理:引入 LifecycleState,拆分启动与关闭流程
- 删除未使用的 TaskExecuteFacadeService 与 TaskSubmitFacadeService

Closes #201

Co-authored-by: AAEE86 <ppk0227@hotmail.com>
This commit is contained in:
fawney19
2026-03-07 15:33:29 +08:00
co-authored by AAEE86
parent 4cd6e0d10f
commit 239238fe47
47 changed files with 3988 additions and 2447 deletions
+245 -3
View File
@@ -6,8 +6,13 @@ import httpx
import pytest
from src.config.settings import config
from src.services.candidate.submit import AllCandidatesFailedError, UpstreamClientRequestError
from src.services.candidate.submit import (
AllCandidatesFailedError,
SubmitOutcome,
UpstreamClientRequestError,
)
from src.services.task.service import TaskService
from src.services.task.submit.outcome_builder import SubmitPayloadParseResult
def _make_candidate(
@@ -356,7 +361,7 @@ async def test_submit_with_failover_applies_pool_reorder_before_submit(
reordered = [pool_candidate_b, pool_candidate_a]
apply_pool_reorder = AsyncMock(return_value=(reordered, []))
monkeypatch.setattr(svc, "_apply_pool_reorder", apply_pool_reorder)
monkeypatch.setattr(svc._submit_ops, "_apply_pool_reorder", apply_pool_reorder) # type: ignore[attr-defined]
submit = AsyncMock(return_value=httpx.Response(200, json={"id": "task-pooled"}))
body = {"session_id": "sid-123"}
@@ -379,8 +384,245 @@ async def test_submit_with_failover_applies_pool_reorder_before_submit(
assert outcome.candidate.key.id == "k-b"
assert submit.await_count == 1
fetch_candidates.assert_awaited_once()
assert fetch_candidates.await_args.kwargs.get("request_body") == body
await_args = fetch_candidates.await_args
assert await_args is not None
assert await_args.kwargs.get("request_body") == body
apply_pool_reorder.assert_awaited_once_with(
[pool_candidate_a, pool_candidate_b],
request_body=body,
)
@pytest.mark.asyncio
async def test_submit_with_failover_orchestrates_prepare_and_execute_layers() -> None:
db = MagicMock()
svc = TaskService(db)
candidate = _make_candidate(provider_id="p2", endpoint_id="e2", key_id="k2")
prepared = SimpleNamespace(candidates=[candidate], record_map={(0, 0): "rc-1"})
outcome = SubmitOutcome(
candidate=candidate, # type: ignore[arg-type]
candidate_keys=[{"index": 0, "provider_id": "p2", "selected": True}],
external_task_id="task-layered",
rule_lookup=None,
upstream_payload={"id": "task-layered"},
upstream_headers={"x-test": "1"},
upstream_status_code=200,
)
svc._submit_ops._prepare_ops.prepare_candidates = AsyncMock( # type: ignore[attr-defined, method-assign]
return_value=prepared
)
svc._submit_ops._execute_ops.execute_submit_loop = AsyncMock( # type: ignore[attr-defined, method-assign]
return_value=outcome
)
result = await svc.submit_with_failover(
api_format="openai:video",
model_name="sora",
affinity_key="a1",
user_api_key=MagicMock(user_id="u1"),
request_id="rid-1",
task_type="video",
submit_func=AsyncMock(),
extract_external_task_id=lambda payload: payload.get("id"),
supported_auth_types={"api_key"},
allow_format_conversion=False,
max_candidates=10,
)
assert result.external_task_id == "task-layered"
svc._submit_ops._prepare_ops.prepare_candidates.assert_awaited_once() # type: ignore[attr-defined]
svc._submit_ops._execute_ops.execute_submit_loop.assert_awaited_once() # type: ignore[attr-defined]
@pytest.mark.asyncio
async def test_submit_with_failover_execute_layer_orchestrates_filter_and_attempt() -> None:
db = MagicMock()
svc = TaskService(db)
candidate = _make_candidate(provider_id="p2", endpoint_id="e2", key_id="k2")
prepared = SimpleNamespace(candidates=[candidate], record_map={(0, 0): "rc-1"})
outcome = SubmitOutcome(
candidate=candidate, # type: ignore[arg-type]
candidate_keys=[{"index": 0, "provider_id": "p2", "selected": True}],
external_task_id="task-inner",
rule_lookup=None,
upstream_payload={"id": "task-inner"},
upstream_headers={"x-test": "1"},
upstream_status_code=200,
)
svc._submit_ops._prepare_ops.prepare_candidates = AsyncMock( # type: ignore[attr-defined, method-assign]
return_value=prepared
)
svc._submit_ops._execute_ops._filter_ops.build_candidate_info = MagicMock( # type: ignore[attr-defined, method-assign]
return_value={
"index": 0,
"provider_id": "p2",
"provider_name": "prov",
"endpoint_id": "e2",
"key_id": "k2",
"key_name": "key",
"auth_type": "api_key",
"priority": 0,
"is_cached": False,
}
)
svc._submit_ops._execute_ops._filter_ops.prepare_candidate_for_attempt = MagicMock( # type: ignore[attr-defined, method-assign]
return_value=SimpleNamespace(record_id="rc-1", rule_lookup=None)
)
svc._submit_ops._execute_ops._attempt_ops.submit_candidate = AsyncMock( # type: ignore[attr-defined, method-assign]
return_value=(outcome, 200)
)
result = await svc.submit_with_failover(
api_format="openai:video",
model_name="sora",
affinity_key="a1",
user_api_key=MagicMock(user_id="u1"),
request_id="rid-2",
task_type="video",
submit_func=AsyncMock(),
extract_external_task_id=lambda payload: payload.get("id"),
supported_auth_types={"api_key"},
allow_format_conversion=False,
max_candidates=10,
)
assert result.external_task_id == "task-inner"
svc._submit_ops._execute_ops._filter_ops.build_candidate_info.assert_called_once() # type: ignore[attr-defined]
svc._submit_ops._execute_ops._filter_ops.prepare_candidate_for_attempt.assert_called_once() # type: ignore[attr-defined]
svc._submit_ops._execute_ops._attempt_ops.submit_candidate.assert_awaited_once() # type: ignore[attr-defined]
@pytest.mark.asyncio
async def test_submit_with_failover_attempt_layer_delegates_to_response_ops() -> None:
db = MagicMock()
svc = TaskService(db)
candidate = _make_candidate(provider_id="p3", endpoint_id="e3", key_id="k3")
prepared = SimpleNamespace(candidates=[candidate], record_map={(0, 0): "rc-3"})
outcome = SubmitOutcome(
candidate=candidate, # type: ignore[arg-type]
candidate_keys=[{"index": 0, "provider_id": "p3", "selected": True}],
external_task_id="task-attempt",
rule_lookup=None,
upstream_payload={"id": "task-attempt"},
upstream_headers={"x-test": "1"},
upstream_status_code=200,
)
svc._submit_ops._prepare_ops.prepare_candidates = AsyncMock( # type: ignore[attr-defined, method-assign]
return_value=prepared
)
svc._submit_ops._execute_ops._filter_ops.build_candidate_info = MagicMock( # type: ignore[attr-defined, method-assign]
return_value={
"index": 0,
"provider_id": "p3",
"provider_name": "prov3",
"endpoint_id": "e3",
"key_id": "k3",
"key_name": "key3",
"auth_type": "api_key",
"priority": 0,
"is_cached": False,
}
)
svc._submit_ops._execute_ops._filter_ops.prepare_candidate_for_attempt = MagicMock( # type: ignore[attr-defined, method-assign]
return_value=SimpleNamespace(record_id="rc-3", rule_lookup=None)
)
svc._submit_ops._execute_ops._attempt_ops._record_ops.mark_pending = MagicMock() # type: ignore[attr-defined, method-assign]
svc._submit_ops._execute_ops._attempt_ops._response_ops.handle_submit_response = MagicMock( # type: ignore[attr-defined, method-assign]
return_value=(outcome, 200)
)
submit = AsyncMock(return_value=httpx.Response(200, json={"id": "task-attempt"}))
result = await svc.submit_with_failover(
api_format="openai:video",
model_name="sora",
affinity_key="a1",
user_api_key=MagicMock(user_id="u1"),
request_id="rid-3",
task_type="video",
submit_func=submit,
extract_external_task_id=lambda payload: payload.get("id"),
supported_auth_types={"api_key"},
allow_format_conversion=False,
max_candidates=10,
)
assert result.external_task_id == "task-attempt"
svc._submit_ops._execute_ops._attempt_ops._record_ops.mark_pending.assert_called_once() # type: ignore[attr-defined]
submit.assert_awaited_once_with(candidate)
svc._submit_ops._execute_ops._attempt_ops._response_ops.handle_submit_response.assert_called_once() # type: ignore[attr-defined]
@pytest.mark.asyncio
async def test_submit_with_failover_response_layer_delegates_to_decider_and_builder() -> None:
db = MagicMock()
svc = TaskService(db)
candidate = _make_candidate(provider_id="p4", endpoint_id="e4", key_id="k4")
prepared = SimpleNamespace(candidates=[candidate], record_map={(0, 0): "rc-4"})
outcome = SubmitOutcome(
candidate=candidate, # type: ignore[arg-type]
candidate_keys=[{"index": 0, "provider_id": "p4", "selected": True}],
external_task_id="task-response",
rule_lookup=None,
upstream_payload={"id": "task-response"},
upstream_headers={"x-test": "1"},
upstream_status_code=200,
)
svc._submit_ops._prepare_ops.prepare_candidates = AsyncMock( # type: ignore[attr-defined, method-assign]
return_value=prepared
)
svc._submit_ops._execute_ops._filter_ops.build_candidate_info = MagicMock( # type: ignore[attr-defined, method-assign]
return_value={
"index": 0,
"provider_id": "p4",
"provider_name": "prov4",
"endpoint_id": "e4",
"key_id": "k4",
"key_name": "key4",
"auth_type": "api_key",
"priority": 0,
"is_cached": False,
}
)
svc._submit_ops._execute_ops._filter_ops.prepare_candidate_for_attempt = MagicMock( # type: ignore[attr-defined, method-assign]
return_value=SimpleNamespace(record_id="rc-4", rule_lookup=None)
)
response_ops = svc._submit_ops._execute_ops._attempt_ops._response_ops # type: ignore[attr-defined]
response_ops._rule_decider.detect_error_stop_pattern = MagicMock(return_value=None) # type: ignore[attr-defined, method-assign]
response_ops._rule_decider.detect_success_failover_pattern = MagicMock(return_value=None) # type: ignore[attr-defined, method-assign]
response_ops._outcome_builder.parse_payload = MagicMock( # type: ignore[attr-defined, method-assign]
return_value=SubmitPayloadParseResult(payload={"id": "task-response"})
)
response_ops._outcome_builder.build_success_outcome = MagicMock( # type: ignore[attr-defined, method-assign]
return_value=outcome
)
submit = AsyncMock(return_value=httpx.Response(200, json={"id": "task-response"}))
result = await svc.submit_with_failover(
api_format="openai:video",
model_name="sora",
affinity_key="a1",
user_api_key=MagicMock(user_id="u1"),
request_id="rid-4",
task_type="video",
submit_func=submit,
extract_external_task_id=lambda payload: payload.get("id"),
supported_auth_types={"api_key"},
allow_format_conversion=False,
max_candidates=10,
)
assert result.external_task_id == "task-response"
response_ops._rule_decider.detect_error_stop_pattern.assert_not_called() # type: ignore[attr-defined]
response_ops._rule_decider.detect_success_failover_pattern.assert_called_once() # type: ignore[attr-defined]
response_ops._outcome_builder.parse_payload.assert_called_once() # type: ignore[attr-defined]
response_ops._outcome_builder.build_success_outcome.assert_called_once() # type: ignore[attr-defined]
+35 -24
View File
@@ -1,16 +1,17 @@
from __future__ import annotations
from __future__ import annotations
from types import SimpleNamespace
from typing import Any, AsyncIterator
from typing import Any, AsyncIterator, cast
from unittest.mock import AsyncMock, MagicMock
import pytest
from src.models.database import Provider, ProviderAPIKey, ProviderEndpoint
from src.services.candidate.failover import FailoverEngine
from src.services.candidate.policy import RetryMode, RetryPolicy, SkipPolicy
from src.services.orchestration.error_classifier import ErrorAction
from src.services.scheduling.schemas import PoolCandidate
from src.services.task.protocol import AttemptKind, AttemptResult
from src.services.orchestration.error_classifier import ErrorAction, ErrorClassifier
from src.services.scheduling.schemas import PoolCandidate, ProviderCandidate
from src.services.task.core.protocol import AttemptKind, AttemptResult
def _make_candidate(
@@ -28,16 +29,22 @@ def _make_candidate(
needs_conversion: bool = False,
provider_max_retries: int | None = None,
provider_config: dict[str, Any] | None = None,
) -> SimpleNamespace:
provider = SimpleNamespace(
id=provider_id,
name=provider_name,
max_retries=provider_max_retries,
config=provider_config,
) -> ProviderCandidate:
provider = cast(
Provider,
SimpleNamespace(
id=provider_id,
name=provider_name,
max_retries=provider_max_retries,
config=provider_config,
),
)
endpoint = SimpleNamespace(id=endpoint_id)
key = SimpleNamespace(id=key_id, name=key_name, auth_type=auth_type, priority=priority)
return SimpleNamespace(
endpoint = cast(ProviderEndpoint, SimpleNamespace(id=endpoint_id))
key = cast(
ProviderAPIKey,
SimpleNamespace(id=key_id, name=key_name, auth_type=auth_type, priority=priority),
)
return ProviderCandidate(
provider=provider,
endpoint=endpoint,
key=key,
@@ -62,10 +69,14 @@ class _StubErrorClassifier:
return self._action
def _stub_classifier(*, action: ErrorAction, client_error: bool = False) -> ErrorClassifier:
return cast(ErrorClassifier, _StubErrorClassifier(action=action, client_error=client_error))
@pytest.mark.asyncio
async def test_failover_engine_success_first_candidate() -> None:
db = MagicMock()
engine = FailoverEngine(db, error_classifier=_StubErrorClassifier(action=ErrorAction.BREAK))
engine = FailoverEngine(db, error_classifier=_stub_classifier(action=ErrorAction.BREAK))
candidates = [_make_candidate(provider_id="p1"), _make_candidate(provider_id="p2")]
@@ -97,7 +108,7 @@ async def test_failover_engine_success_first_candidate() -> None:
@pytest.mark.asyncio
async def test_failover_engine_continue_to_next_candidate_on_error() -> None:
db = MagicMock()
engine = FailoverEngine(db, error_classifier=_StubErrorClassifier(action=ErrorAction.BREAK))
engine = FailoverEngine(db, error_classifier=_stub_classifier(action=ErrorAction.BREAK))
candidates = [_make_candidate(provider_id="p1"), _make_candidate(provider_id="p2")]
@@ -132,7 +143,7 @@ async def test_failover_engine_continue_to_next_candidate_on_error() -> None:
async def test_failover_engine_retry_same_candidate_when_classifier_says_continue() -> None:
db = MagicMock()
# ErrorAction.CONTINUE => retry current candidate (mapped to FailoverAction.RETRY)
engine = FailoverEngine(db, error_classifier=_StubErrorClassifier(action=ErrorAction.CONTINUE))
engine = FailoverEngine(db, error_classifier=_stub_classifier(action=ErrorAction.CONTINUE))
candidates = [_make_candidate(provider_id="p1", is_cached=True, provider_max_retries=2)]
@@ -167,7 +178,7 @@ async def test_failover_engine_continues_when_classifier_raises() -> None:
"""After the 'default failover' change, RAISE no longer stops failover.
All candidates should be attempted."""
db = MagicMock()
engine = FailoverEngine(db, error_classifier=_StubErrorClassifier(action=ErrorAction.RAISE))
engine = FailoverEngine(db, error_classifier=_stub_classifier(action=ErrorAction.RAISE))
candidates = [_make_candidate(provider_id="p1"), _make_candidate(provider_id="p2")]
attempt = AsyncMock(side_effect=RuntimeError("client-ish"))
@@ -200,7 +211,7 @@ async def _empty_stream() -> AsyncIterator[bytes]:
@pytest.mark.asyncio
async def test_failover_engine_stream_probe_wraps_first_chunk() -> None:
db = MagicMock()
engine = FailoverEngine(db, error_classifier=_StubErrorClassifier(action=ErrorAction.BREAK))
engine = FailoverEngine(db, error_classifier=_stub_classifier(action=ErrorAction.BREAK))
candidates = [_make_candidate(provider_id="p1")]
attempt = AsyncMock(
@@ -234,7 +245,7 @@ async def test_failover_engine_stream_probe_wraps_first_chunk() -> None:
@pytest.mark.asyncio
async def test_failover_engine_stream_probe_empty_triggers_failover() -> None:
db = MagicMock()
engine = FailoverEngine(db, error_classifier=_StubErrorClassifier(action=ErrorAction.BREAK))
engine = FailoverEngine(db, error_classifier=_stub_classifier(action=ErrorAction.BREAK))
candidates = [_make_candidate(provider_id="p1"), _make_candidate(provider_id="p2")]
@@ -273,7 +284,7 @@ async def test_failover_engine_pre_expand_marks_unused_slots_on_success(
monkeypatch: pytest.MonkeyPatch,
) -> None:
db = MagicMock()
engine = FailoverEngine(db, error_classifier=_StubErrorClassifier(action=ErrorAction.BREAK))
engine = FailoverEngine(db, error_classifier=_stub_classifier(action=ErrorAction.BREAK))
# patch low-level record updater to observe unused marking
engine._update_record = MagicMock() # type: ignore[method-assign]
@@ -333,7 +344,7 @@ class _HttpError(Exception):
async def test_error_stop_pattern_with_matching_status_code_stops_failover() -> None:
"""When status_codes is set and matches, failover should stop."""
db = MagicMock()
engine = FailoverEngine(db, error_classifier=_StubErrorClassifier(action=ErrorAction.BREAK))
engine = FailoverEngine(db, error_classifier=_stub_classifier(action=ErrorAction.BREAK))
config = {
"failover_rules": {
@@ -376,7 +387,7 @@ async def test_error_stop_pattern_with_matching_status_code_stops_failover() ->
async def test_error_stop_pattern_with_non_matching_status_code_continues() -> None:
"""When status_codes is set but doesn't match, the rule is skipped and failover continues."""
db = MagicMock()
engine = FailoverEngine(db, error_classifier=_StubErrorClassifier(action=ErrorAction.BREAK))
engine = FailoverEngine(db, error_classifier=_stub_classifier(action=ErrorAction.BREAK))
config = {
"failover_rules": {
@@ -420,7 +431,7 @@ async def test_error_stop_pattern_with_non_matching_status_code_continues() -> N
async def test_error_stop_pattern_without_status_codes_matches_any() -> None:
"""When status_codes is not set, the rule matches any status code (existing behavior)."""
db = MagicMock()
engine = FailoverEngine(db, error_classifier=_StubErrorClassifier(action=ErrorAction.BREAK))
engine = FailoverEngine(db, error_classifier=_stub_classifier(action=ErrorAction.BREAK))
config = {
"failover_rules": {
@@ -0,0 +1,132 @@
from __future__ import annotations
from types import SimpleNamespace
from unittest.mock import MagicMock
import pytest
from src.services.candidate.policy import FailoverAction
from src.services.task.execute.exception_classification import (
CandidateErrorAction,
classify_candidate_error_action,
)
from src.services.task.execute.state_transition import (
SyncExecutionState,
resolve_execution_error_transition,
)
def _make_candidate() -> SimpleNamespace:
return SimpleNamespace(
provider=SimpleNamespace(id="p1", name="provider-1"),
endpoint=SimpleNamespace(id="e1"),
key=SimpleNamespace(id="k1"),
)
@pytest.mark.parametrize(
("raw_action", "expected"),
[
("continue", CandidateErrorAction.RETRY_CURRENT),
("break", CandidateErrorAction.NEXT_CANDIDATE),
("raise", CandidateErrorAction.RAISE_ERROR),
("unexpected", CandidateErrorAction.NEXT_CANDIDATE),
(None, CandidateErrorAction.NEXT_CANDIDATE),
],
)
def test_classify_candidate_error_action(
raw_action: str | None, expected: CandidateErrorAction
) -> None:
assert classify_candidate_error_action(raw_action) == expected
def test_resolve_execution_error_transition_retry_and_consume_rectify_flag() -> None:
request_body_ref = {"_rectified_this_turn": True}
state = SyncExecutionState(
candidate_record_map={},
request_body_ref=request_body_ref,
)
transition = resolve_execution_error_transition(
action=CandidateErrorAction.RETRY_CURRENT,
state=state,
max_retries_for_candidate=2,
retry_index=1,
)
assert transition.failover_action == FailoverAction.RETRY
assert transition.max_retries == 3
assert request_body_ref["_rectified_this_turn"] is False
def test_resolve_execution_error_transition_next_candidate() -> None:
state = SyncExecutionState(
candidate_record_map={},
request_body_ref={"_rectified_this_turn": True},
)
transition = resolve_execution_error_transition(
action=CandidateErrorAction.NEXT_CANDIDATE,
state=state,
max_retries_for_candidate=2,
retry_index=0,
)
assert transition.failover_action == FailoverAction.CONTINUE
assert transition.max_retries is None
assert state.request_body_ref == {"_rectified_this_turn": True}
def test_sync_execution_state_resolve_candidate_record_id_fallback() -> None:
state = SyncExecutionState(
candidate_record_map={(2, 0): "r20"},
request_body_ref=None,
)
assert state.resolve_candidate_record_id(candidate_index=2, record_id=None) == "r20"
assert state.resolve_candidate_record_id(candidate_index=2, record_id="r22") == "r22"
def test_sync_execution_state_raise_classified_error_uses_last_error() -> None:
err = ValueError("boom")
candidate = _make_candidate()
state = SyncExecutionState(
candidate_record_map={},
request_body_ref=None,
last_error=err,
last_candidate=candidate,
)
failure_ops = MagicMock()
with pytest.raises(ValueError, match="boom"):
state.raise_classified_error(
fallback_error=RuntimeError("fallback"),
failure_ops=failure_ops,
model_name="gpt-4.1",
api_format="openai_chat",
)
failure_ops.attach_metadata_to_error.assert_called_once_with(
err,
candidate,
"gpt-4.1",
"openai_chat",
)
def test_sync_execution_state_raise_classified_error_fallback_error() -> None:
state = SyncExecutionState(
candidate_record_map={},
request_body_ref=None,
)
failure_ops = MagicMock()
with pytest.raises(RuntimeError, match="fallback"):
state.raise_classified_error(
fallback_error=RuntimeError("fallback"),
failure_ops=failure_ops,
model_name="gpt-4.1",
api_format="openai_chat",
)
failure_ops.attach_metadata_to_error.assert_not_called()
+118
View File
@@ -0,0 +1,118 @@
from __future__ import annotations
from types import SimpleNamespace
from typing import Any, cast
import pytest
from src.models.database import Provider, ProviderAPIKey, ProviderEndpoint
from src.services.scheduling.schemas import PoolCandidate, ProviderCandidate
from src.services.task.execute.pool import TaskPoolOperationsService
def _provider(provider_id: str) -> Provider:
return cast(Provider, SimpleNamespace(id=provider_id))
def _endpoint(endpoint_id: str) -> ProviderEndpoint:
return cast(ProviderEndpoint, SimpleNamespace(id=endpoint_id))
def _key(key_id: str) -> ProviderAPIKey:
return cast(ProviderAPIKey, SimpleNamespace(id=key_id))
def test_extract_session_uuid_returns_none_for_non_dict_request_body() -> None:
svc = TaskPoolOperationsService()
assert svc.extract_session_uuid("openai", None) is None
assert svc.extract_session_uuid("openai", cast(Any, "not-dict")) is None
def test_extract_session_uuid_uses_pool_hook(monkeypatch: pytest.MonkeyPatch) -> None:
from src.services.provider.pool import hooks
hook = SimpleNamespace(extract_session_uuid=lambda body: f"sid:{body.get('session')}")
def _get_pool_hook(_provider_type: str) -> Any:
return hook
monkeypatch.setattr(hooks, "get_pool_hook", _get_pool_hook)
svc = TaskPoolOperationsService()
session_id = svc.extract_session_uuid("claude_code", {"session": "abc"})
assert session_id == "sid:abc"
def test_expand_pool_candidates_for_async_submit_keeps_non_pool_candidate() -> None:
provider = _provider("p1")
endpoint = _endpoint("e1")
key = _key("k1")
candidate = ProviderCandidate(provider=provider, endpoint=endpoint, key=key)
svc = TaskPoolOperationsService()
expanded = svc.expand_pool_candidates_for_async_submit([candidate])
assert len(expanded) == 1
assert expanded[0] is candidate
def test_expand_pool_candidates_for_async_submit_expands_pool_keys() -> None:
provider = _provider("p1")
endpoint = _endpoint("e1")
key = _key("k0")
pool_key_1 = cast(
ProviderAPIKey,
SimpleNamespace(
id="k1",
_pool_skipped=False,
_pool_mapping_matched_model="mapped-model",
_pool_extra_data={"source": "warm"},
),
)
pool_key_2 = cast(
ProviderAPIKey,
SimpleNamespace(
id="k2",
_pool_skipped=True,
_pool_skip_reason="cooldown",
_pool_extra_data={"reason_code": "429"},
),
)
pool_candidate = PoolCandidate(
provider=provider,
endpoint=endpoint,
key=key,
is_cached=True,
is_skipped=False,
skip_reason=None,
mapping_matched_model="fallback-model",
needs_conversion=True,
provider_api_format="openai:chat",
output_limit=1024,
capability_miss_count=1,
pool_keys=[pool_key_1, pool_key_2],
)
svc = TaskPoolOperationsService()
expanded = svc.expand_pool_candidates_for_async_submit([pool_candidate])
assert len(expanded) == 2
first, second = expanded
assert first.key is pool_key_1
assert first.is_skipped is False
assert first.mapping_matched_model == "mapped-model"
first_extra = getattr(first, "_pool_extra_data")
assert first_extra["pool_group_id"] == "p1"
assert first_extra["pool_key_index"] == 0
assert first_extra["source"] == "warm"
assert second.key is pool_key_2
assert second.is_skipped is True
assert second.skip_reason == "cooldown"
assert second.mapping_matched_model == "fallback-model"
second_extra = getattr(second, "_pool_extra_data")
assert second_extra["pool_group_id"] == "p1"
assert second_extra["pool_key_index"] == 1
assert second_extra["reason_code"] == "429"
@@ -1,4 +1,4 @@
from __future__ import annotations
from __future__ import annotations
from types import SimpleNamespace
from unittest.mock import AsyncMock, MagicMock
@@ -8,10 +8,9 @@ import pytest
from src.core.exceptions import EmbeddedErrorException
from src.services.candidate.schema import CandidateKey
from src.services.candidate.submit import SubmitOutcome
from src.services.request.executor import ExecutionContext, ExecutionError
from src.services.task import service as task_service_module
from src.services.task.context import TaskMode
from src.services.task.protocol import AttemptKind
from src.services.task.core.context import TaskMode
from src.services.task.core.protocol import AttemptKind
from src.services.task.service import pool_on_error
from src.services.task.service import TaskService
@@ -51,7 +50,7 @@ async def test_task_service_execute_async_returns_execution_result() -> None:
)
svc.submit_with_failover = AsyncMock(return_value=outcome) # type: ignore[method-assign]
svc._recorder.get_candidate_keys = MagicMock( # type: ignore[attr-defined, method-assign]
svc._execute_facade_ops._get_candidate_keys = MagicMock( # type: ignore[attr-defined, method-assign]
return_value=[
CandidateKey(candidate_index=0, retry_index=0, status="success", provider_id="p1")
]
@@ -82,7 +81,7 @@ async def test_task_service_execute_sync_passes_request_headers_and_body() -> No
db = MagicMock()
svc = TaskService(db)
sentinel_result = object()
svc._execute_sync_unified = AsyncMock( # type: ignore[method-assign]
svc._sync_ops.execute_sync_unified = AsyncMock( # type: ignore[attr-defined, method-assign]
return_value=sentinel_result
)
@@ -105,95 +104,92 @@ async def test_task_service_execute_sync_passes_request_headers_and_body() -> No
)
assert result is sentinel_result
svc._execute_sync_unified.assert_awaited_once() # type: ignore[attr-defined]
kwargs = svc._execute_sync_unified.await_args.kwargs # type: ignore[attr-defined, union-attr]
svc._sync_ops.execute_sync_unified.assert_awaited_once() # type: ignore[attr-defined]
kwargs = svc._sync_ops.execute_sync_unified.await_args.kwargs # type: ignore[attr-defined, union-attr]
assert kwargs["request_headers"] == request_headers
assert kwargs["request_body"] == request_body
assert kwargs["request_body_ref"] == request_body_ref
@pytest.mark.asyncio
@pytest.mark.parametrize(
("is_client_error", "retry_index", "max_retries_for_candidate", "expected_action"),
[
(True, 0, 1, "break"),
(False, 0, 2, "continue"),
],
)
async def test_task_service_embedded_error_branch_applies_pool_health_policy(
monkeypatch: pytest.MonkeyPatch,
is_client_error: bool,
retry_index: int,
max_retries_for_candidate: int,
expected_action: str,
) -> None:
db = MagicMock()
svc = TaskService(db)
async def test_task_service_execute_delegates_to_execute_facade_ops() -> None:
svc = TaskService(MagicMock())
sentinel = object()
svc._execute_facade_ops.execute = AsyncMock(return_value=sentinel) # type: ignore[attr-defined, method-assign]
monkeypatch.setattr(
task_service_module.RequestCandidateService, "mark_candidate_failed", MagicMock()
)
monkeypatch.setattr(
"src.services.proxy_node.resolver.resolve_effective_proxy",
lambda provider_proxy, key_proxy: provider_proxy or key_proxy,
)
monkeypatch.setattr("src.services.proxy_node.resolver.resolve_proxy_info", lambda _proxy: None)
pool_on_error = AsyncMock()
monkeypatch.setattr(svc, "_pool_on_error", pool_on_error)
candidate = SimpleNamespace(
provider=SimpleNamespace(id="p1", name="prov", proxy=None),
endpoint=SimpleNamespace(id="e1"),
key=SimpleNamespace(id="k1", proxy=None),
)
cause = EmbeddedErrorException(
provider_name="prov",
error_code=429,
error_message="usage_limit_reached",
error_status="RESOURCE_EXHAUSTED",
)
context = ExecutionContext(
candidate_id="cid-1",
candidate_index=0,
provider_id="p1",
endpoint_id="e1",
key_id="k1",
user_id=None,
api_key_id=None,
is_cached_user=False,
elapsed_ms=12,
concurrent_requests=3,
)
exec_err = ExecutionError(cause, context)
classifier = SimpleNamespace(is_client_error=lambda _text: is_client_error)
action = await svc._handle_candidate_error(
exec_err=exec_err,
candidate=candidate,
candidate_record_id="cand-1",
retry_index=retry_index,
max_retries_for_candidate=max_retries_for_candidate,
affinity_key="provider-test:p1",
result = await svc.execute(
task_type="chat",
task_mode=TaskMode.SYNC,
api_format="openai:chat",
global_model_id="gpt-4o-mini",
request_id="req-1",
attempt=1,
max_attempts=3,
error_classifier=classifier,
model_name="m",
user_api_key=MagicMock(id="u", user_id="user"),
request_func=AsyncMock(),
request_id="rid",
)
assert action == expected_action
pool_on_error.assert_awaited_once_with(candidate.provider, candidate.key, 429, cause)
assert result is sentinel
svc._execute_facade_ops.execute.assert_awaited_once() # type: ignore[attr-defined]
@pytest.mark.asyncio
async def test_task_service_pool_on_error_uses_embedded_error_message_fallback(
async def test_task_service_submit_with_failover_delegates_to_submit_facade_ops() -> None:
svc = TaskService(MagicMock())
sentinel = object()
svc._submit_facade_ops.submit_with_failover = AsyncMock( # type: ignore[attr-defined, method-assign]
return_value=sentinel
)
result = await svc.submit_with_failover(
api_format="openai:video",
model_name="sora",
affinity_key="a1",
user_api_key=MagicMock(id="u", user_id="user"),
request_id="rid",
task_type="video",
submit_func=AsyncMock(),
extract_external_task_id=lambda payload: payload.get("id"),
)
assert result is sentinel
svc._submit_facade_ops.submit_with_failover.assert_awaited_once() # type: ignore[attr-defined]
@pytest.mark.asyncio
async def test_task_service_poll_delegates_to_video_facade_ops() -> None:
svc = TaskService(MagicMock())
sentinel = object()
svc._video_facade_ops.poll = AsyncMock(return_value=sentinel) # type: ignore[attr-defined, method-assign]
result = await svc.poll("task-1", user_id="user-1")
assert result is sentinel
svc._video_facade_ops.poll.assert_awaited_once_with("task-1", user_id="user-1") # type: ignore[attr-defined]
@pytest.mark.asyncio
async def test_task_service_cancel_delegates_to_video_facade_ops() -> None:
svc = TaskService(MagicMock())
sentinel = object()
svc._video_facade_ops.cancel = AsyncMock(return_value=sentinel) # type: ignore[attr-defined, method-assign]
result = await svc.cancel(
"task-1",
user_id="user-1",
original_headers={"x-test": "1"},
)
assert result is sentinel
svc._video_facade_ops.cancel.assert_awaited_once_with( # type: ignore[attr-defined]
"task-1",
user_id="user-1",
original_headers={"x-test": "1"},
)
@pytest.mark.asyncio
async def test_pool_on_error_uses_embedded_error_message_fallback(
monkeypatch: pytest.MonkeyPatch,
) -> None:
db = MagicMock()
svc = TaskService(db)
parsed_pool_cfg = object()
apply_health_policy = AsyncMock()
monkeypatch.setattr(
@@ -212,7 +208,7 @@ async def test_task_service_pool_on_error_uses_embedded_error_message_fallback(
error_message="usage_limit_reached",
)
await svc._pool_on_error(provider, key, 429, cause)
await pool_on_error(provider, key, 429, cause)
apply_health_policy.assert_awaited_once()
kwargs = apply_health_policy.await_args.kwargs
@@ -0,0 +1,44 @@
from __future__ import annotations
from types import SimpleNamespace
from unittest.mock import AsyncMock, MagicMock
import pytest
from fastapi import HTTPException
from src.services.task.core.exceptions import TaskNotFoundError
from src.services.task.video.operations import VideoTaskOperationsService
@pytest.mark.asyncio
async def test_video_ops_cancel_delegates_to_cancel_ops() -> None:
svc = VideoTaskOperationsService(MagicMock())
task = SimpleNamespace(id="task-1")
svc._get_video_task_for_user = MagicMock(return_value=task) # type: ignore[attr-defined, method-assign]
svc._cancel_ops.cancel_task = AsyncMock(return_value={"ok": True}) # type: ignore[attr-defined, method-assign]
result = await svc.cancel(
"task-1",
user_id="user-1",
original_headers={"x-test": "1"},
)
assert result == {"ok": True}
svc._cancel_ops.cancel_task.assert_awaited_once_with( # type: ignore[attr-defined]
task=task,
task_id="task-1",
original_headers={"x-test": "1"},
)
@pytest.mark.asyncio
async def test_video_ops_cancel_maps_not_found_to_http_404() -> None:
svc = VideoTaskOperationsService(MagicMock())
svc._get_video_task_for_user = MagicMock(side_effect=TaskNotFoundError("missing")) # type: ignore[attr-defined, method-assign]
with pytest.raises(HTTPException) as excinfo:
await svc.cancel("missing", user_id="user-1")
assert excinfo.value.status_code == 404
assert excinfo.value.detail == "Video task not found"
+14 -9
View File
@@ -1,10 +1,12 @@
from types import SimpleNamespace
from types import SimpleNamespace
from typing import cast
from unittest.mock import AsyncMock, MagicMock
import pytest
from src.core.api_format.conversion.internal_video import InternalVideoPollResult, VideoStatus
from src.services.task.impl.video_poller import VideoTaskPollerAdapter
from src.models.database import VideoTask
from src.services.task.video.poller_adapter import VideoTaskPollerAdapter
@pytest.mark.asyncio
@@ -13,11 +15,14 @@ async def test_poll_task_status_routes_gemini_video_to_gemini(
) -> None:
adapter = VideoTaskPollerAdapter()
task = SimpleNamespace(
endpoint_id="e1",
key_id="k1",
provider_api_format="gemini:video",
external_task_id="operations/123",
task = cast(
VideoTask,
SimpleNamespace(
endpoint_id="e1",
key_id="k1",
provider_api_format="gemini:video",
external_task_id="operations/123",
),
)
endpoint = SimpleNamespace(id="e1", base_url="https://example.com", api_format="gemini:video")
key = SimpleNamespace(id="k1", api_key="enc")
@@ -25,12 +30,12 @@ async def test_poll_task_status_routes_gemini_video_to_gemini(
monkeypatch.setattr(adapter, "_get_endpoint", lambda _db, _id: endpoint)
monkeypatch.setattr(adapter, "_get_key", lambda _db, _id: key)
monkeypatch.setattr(
"src.services.task.impl.video_poller.crypto_service.decrypt", lambda _v: "decrypted"
"src.services.task.video.poller_adapter.crypto_service.decrypt", lambda _v: "decrypted"
)
auth_info = SimpleNamespace(auth_header="authorization", auth_value="Bearer x")
monkeypatch.setattr(
"src.services.task.impl.video_poller.get_provider_auth",
"src.services.task.video.poller_adapter.get_provider_auth",
AsyncMock(return_value=auth_info),
)