Files
SmartUp/backend/test_website_client.py
T

424 lines
14 KiB
Python

from decimal import Decimal
import httpx
import pytest
from app.services.website_client import (
WebsiteError,
_friendly_connection_error,
_friendly_http_error,
calculate_target_rate,
normalize_groups,
)
# ——— normalize_groups ———
def test_normalize_groups_unwraps_sub2api_paginated_response():
groups = normalize_groups({
"code": 0,
"message": "success",
"data": {
"items": [
{"id": "codex-free", "name": "codex-free", "rate_multiplier": 1},
{"id": "my-plus", "name": "my-plus", "rate_multiplier": "1.5"},
{"id": "deepseek", "name": "deepseek"},
],
"total": 3,
"page": 1,
},
})
assert [group["id"] for group in groups] == ["codex-free", "my-plus", "deepseek"]
assert groups[0]["rate_multiplier"] == "1.00"
assert groups[1]["rate_multiplier"] == "1.50"
assert groups[2]["rate_multiplier"] is None
def test_normalize_groups_unwraps_wrapped_list_response():
groups = normalize_groups({
"code": 0,
"message": "success",
"data": [
{"id": "default", "name": "Default", "rateMultiplier": "2.0"},
],
})
assert groups == [{
"id": "default",
"name": "Default",
"rate_multiplier": "2.00",
"description": None,
"raw": {"id": "default", "name": "Default", "rateMultiplier": "2.0"},
}]
def test_normalize_groups_unwraps_groups_key_response():
groups = normalize_groups({
"data": {
"groups": [
{"group_id": "vip", "group_name": "VIP", "ratio": "0.75"},
],
},
})
assert groups[0]["id"] == "vip"
assert groups[0]["name"] == "VIP"
assert groups[0]["rate_multiplier"] == "0.75"
def test_normalize_groups_keeps_string_list_compatibility():
groups = normalize_groups(["free", "paid"])
assert [group["id"] for group in groups] == ["free", "paid"]
assert groups[0]["raw"] == {"id": "free", "name": "free"}
def test_normalize_groups_keeps_plain_dict_mapping_compatibility():
groups = normalize_groups({
"free": {"id": "free", "name": "Free", "rate_multiplier": "1.00"},
"paid": {"id": "paid", "name": "Paid"},
})
assert [group["id"] for group in groups] == ["free", "paid"]
assert groups[0]["rate_multiplier"] == "1.00"
def test_update_group_rate_writes_rounded_two_decimal_number():
from app.services.website_client import Sub2ApiWebsiteClient
requests = []
def handler(req: httpx.Request) -> httpx.Response:
requests.append(req)
return httpx.Response(200, json={"ok": True}, request=req)
client = Sub2ApiWebsiteClient(
base_url="https://target.example",
api_prefix="/api/v1/admin",
auth_type="api_key",
auth_config={"key": "admin-key"},
)
client._client = httpx.Client(transport=httpx.MockTransport(handler))
try:
client.update_group_rate("/groups/{id}", "vip", Decimal("2.2"))
client.update_group_rate("/groups/{id}", "free", Decimal("1"))
client.update_group_rate("/groups/{id}", "small", Decimal("0.1275"))
finally:
client.close()
assert requests[0].url.path == "/api/v1/admin/groups/vip"
assert requests[0].read() == b'{"rate_multiplier":2.2}'
assert requests[1].url.path == "/api/v1/admin/groups/free"
assert requests[1].read() == b'{"rate_multiplier":1.0}'
assert requests[2].url.path == "/api/v1/admin/groups/small"
assert requests[2].read() == b'{"rate_multiplier":0.13}'
def test_calculate_target_rate_rounds_to_two_decimal_places():
assert calculate_target_rate(["0.1275"], 0, "max_plus_percent") == Decimal("0.13")
# ——— _get_account_ids / account_exists ———
def test_get_account_ids_flat_list():
from app.services.website_client import Sub2ApiWebsiteClient
ids = Sub2ApiWebsiteClient._unwrap_list([
{"id": 1, "name": "a"}, {"id": 2, "name": "b"},
])
assert ids == [{"id": 1, "name": "a"}, {"id": 2, "name": "b"}]
def test_get_account_ids_top_level_items():
from app.services.website_client import Sub2ApiWebsiteClient
ids = Sub2ApiWebsiteClient._unwrap_list({"items": [
{"id": "k1"}, {"id": "k2"},
]})
assert ids == [{"id": "k1"}, {"id": "k2"}]
def test_get_account_ids_nested_data_items():
from app.services.website_client import Sub2ApiWebsiteClient
ids = Sub2ApiWebsiteClient._unwrap_list({"data": {"items": [
{"id": "a1", "name": "Alpha"},
{"id": "a2", "name": "Beta"},
]}})
assert ids == [{"id": "a1", "name": "Alpha"}, {"id": "a2", "name": "Beta"}]
def test_get_account_ids_nested_data_empty():
from app.services.website_client import Sub2ApiWebsiteClient
ids = Sub2ApiWebsiteClient._unwrap_list({"data": {"items": []}})
assert ids == []
def test_get_account_ids_unexpected_format():
from app.services.website_client import Sub2ApiWebsiteClient
ids = Sub2ApiWebsiteClient._unwrap_list({"error": "not found"})
assert ids is None
# ——— 友好错误提示 ———
def _make_response(status_code: int, path: str = "/groups") -> httpx.Response:
"""创建模拟的 httpx.Response 用于错误测试。"""
req = httpx.Request("GET", f"http://target.local/api/v1{path}")
resp = httpx.Response(status_code, request=req)
return resp
def test_friendly_http_401():
resp = _make_response(401)
exc = httpx.HTTPStatusError("401", request=resp.request, response=resp)
msg = _friendly_http_error(exc)
assert "认证失败" in msg
assert "API Key" in msg
assert "http://" not in msg
def test_friendly_http_403():
resp = _make_response(403)
exc = httpx.HTTPStatusError("403", request=resp.request, response=resp)
msg = _friendly_http_error(exc)
assert "权限不足" in msg
def test_friendly_http_404():
resp = _make_response(404, path="/wrong-path")
exc = httpx.HTTPStatusError("404", request=resp.request, response=resp)
msg = _friendly_http_error(exc)
assert "接口不存在" in msg
assert "/wrong-path" in msg
# 不包含完整 URL / MDN 链接
assert "http://" not in msg
assert "MDN" not in msg
def test_friendly_http_500():
resp = _make_response(502)
exc = httpx.HTTPStatusError("502", request=resp.request, response=resp)
msg = _friendly_http_error(exc)
assert "服务异常" in msg
def test_friendly_connect_error():
exc = httpx.ConnectError("Connection refused")
msg = _friendly_connection_error(exc)
assert "无法连接" in msg
def test_friendly_timeout_error():
exc = httpx.TimeoutException("Timed out")
msg = _friendly_connection_error(exc)
assert "请求超时" in msg
# ——— get_groups fallback(通过 mock httpx client 触发真实 _request 错误转换) ———
def _mock_httpx_request(status_code: int, path: str = "/groups"):
"""返回一个 mock 的 httpx.Client.request,直接抛 HTTPStatusError。"""
def request(self, method, url, **kwargs):
resp = _make_response(status_code, path=path)
raise httpx.HTTPStatusError(f"{status_code} {path}", request=resp.request, response=resp)
return request
def test_get_groups_401_returns_friendly_auth_error(monkeypatch):
from app.services.website_client import Sub2ApiWebsiteClient
monkeypatch.setattr(httpx.Client, "request", _mock_httpx_request(401))
client = Sub2ApiWebsiteClient("http://target.local", "api/v1", "api_key", {"key": "bad"})
with pytest.raises(WebsiteError) as excinfo:
client.get_groups("/groups")
msg = str(excinfo.value)
assert "认证失败" in msg
assert "http://" not in msg
assert "MDN" not in msg
def test_get_groups_404_fallback_succeeds(monkeypatch):
from app.services.website_client import Sub2ApiWebsiteClient
call_count = 0
def request_fallback(self, method, url, **kwargs):
nonlocal call_count
call_count += 1
if call_count == 1:
resp = _make_response(404)
raise httpx.HTTPStatusError("404", request=resp.request, response=resp)
req = httpx.Request("GET", url)
return httpx.Response(200, json={"data": [{"id": "default", "name": "Default", "rate_multiplier": "1"}]}, request=req)
monkeypatch.setattr(httpx.Client, "request", request_fallback)
client = Sub2ApiWebsiteClient("http://target.local", "api/v1", "api_key", {"key": "ok"})
groups = client.get_groups("/groups")
assert len(groups) == 1
assert groups[0]["id"] == "default"
assert call_count == 2
def test_get_groups_all_404_no_raw_url(monkeypatch):
from app.services.website_client import Sub2ApiWebsiteClient
monkeypatch.setattr(httpx.Client, "request", _mock_httpx_request(404))
client = Sub2ApiWebsiteClient("http://target.local", "api/v1", "api_key", {"key": "ok"})
with pytest.raises(WebsiteError) as excinfo:
client.get_groups("/groups")
msg = str(excinfo.value)
assert "接口不存在" in msg
assert "http://" not in msg
assert "MDN" not in msg
def test_list_accounts_pagination():
from app.services.website_client import Sub2ApiWebsiteClient
requests_called = []
def handler(req: httpx.Request) -> httpx.Response:
requests_called.append(req)
url_str = str(req.url)
if "page=1" in url_str:
return httpx.Response(200, json={
"total": 3,
"page": 1,
"page_size": 2,
"pages": 2,
"data": [
{"id": "a1", "name": "acc1"},
{"id": "a2", "name": "acc2"},
]
}, request=req)
elif "page=2" in url_str:
return httpx.Response(200, json={
"total": 3,
"page": 2,
"page_size": 2,
"pages": 2,
"data": [
{"id": "a3", "name": "acc3"},
]
}, request=req)
return httpx.Response(404, json={"error": "Not Found"}, request=req)
client = Sub2ApiWebsiteClient(
base_url="https://target.example",
api_prefix="/api/v1/admin",
auth_type="api_key",
auth_config={"key": "admin-key"},
)
client._client = httpx.Client(transport=httpx.MockTransport(handler))
try:
accounts = client.list_accounts("/accounts")
assert accounts == [
{"id": "a1", "name": "acc1"},
{"id": "a2", "name": "acc2"},
{"id": "a3", "name": "acc3"},
]
assert len(requests_called) == 2
assert "page_size=100" in str(requests_called[0].url)
assert "page=1" in str(requests_called[0].url)
assert "page=2" in str(requests_called[1].url)
finally:
client.close()
def test_list_accounts_pagination_single_large_page_does_not_refetch():
from app.services.website_client import Sub2ApiWebsiteClient
requests_called = []
def handler(req: httpx.Request) -> httpx.Response:
requests_called.append(req)
return httpx.Response(200, json={
"total": 2,
"page": 1,
"page_size": 100,
"pages": 1,
"data": [
{"id": "a1", "name": "acc1"},
{"id": "a2", "name": "acc2"},
],
}, request=req)
client = Sub2ApiWebsiteClient(
base_url="https://target.example",
api_prefix="/api/v1/admin",
auth_type="api_key",
auth_config={"key": "admin-key"},
)
client._client = httpx.Client(transport=httpx.MockTransport(handler))
try:
accounts = client.list_accounts("/accounts")
assert accounts == [
{"id": "a1", "name": "acc1"},
{"id": "a2", "name": "acc2"},
]
assert len(requests_called) == 1
assert "page_size=100" in str(requests_called[0].url)
finally:
client.close()
def test_list_accounts_pagination_returns_none_when_any_page_fails():
from app.services.website_client import Sub2ApiWebsiteClient
requests_called = []
def handler(req: httpx.Request) -> httpx.Response:
requests_called.append(req)
if "page=1" in str(req.url):
return httpx.Response(200, json={
"total": 3,
"page": 1,
"page_size": 2,
"pages": 2,
"data": [
{"id": "a1", "name": "acc1"},
{"id": "a2", "name": "acc2"},
],
}, request=req)
return httpx.Response(500, json={"error": "failed"}, request=req)
client = Sub2ApiWebsiteClient(
base_url="https://target.example",
api_prefix="/api/v1/admin",
auth_type="api_key",
auth_config={"key": "admin-key"},
)
client._client = httpx.Client(transport=httpx.MockTransport(handler))
try:
assert client.list_accounts("/accounts") is None
assert len(requests_called) == 2
finally:
client.close()
def test_calculate_target_rate_priority_weighted():
# 0.18、0.25 应当至少满足 min * 1.50 (0.27) 和 max * 1.10 (0.275),向上取整为 0.28,忽略 percent 输入
assert calculate_target_rate(["0.18", "0.25"], 50, "priority_weighted_plus_percent") == Decimal("0.28")
assert calculate_target_rate(["0.18", "0.25"], 0, "priority_weighted_plus_percent") == Decimal("0.28")
# 单来源 0.05 自动得到 0.08 (max(0.05 * 1.50, 0.05 * 1.10) = 0.075 -> ROUND_UP -> 0.08)
assert calculate_target_rate(["0.05"], 0, "priority_weighted_plus_percent") == Decimal("0.08")
# 4 个来源:min=0.18, max=0.36. max(0.27, 0.396) = 0.396 -> ROUND_UP -> 0.40
assert calculate_target_rate(["0.18", "0.24", "0.30", "0.36"], 50, "priority_weighted_plus_percent") == Decimal("0.40")
# 验证保护线与精度保护:最低来源 0.124。max(0.186, 0.1364) = 0.186 -> ROUND_UP -> 0.19
assert calculate_target_rate(["0.124"], 0, "priority_weighted_plus_percent") == Decimal("0.19")
# 无可用正数倍率时报错
with pytest.raises(WebsiteError, match="没有可用的正数上游倍率"):
calculate_target_rate([], 50, "priority_weighted_plus_percent")
# 旧算法结果不变
assert calculate_target_rate(["0.18", "0.25"], 50, "max_plus_percent") == Decimal("0.38")