feat: associate generation calls with devices

This commit is contained in:
QiuSW
2026-07-21 09:40:40 +08:00
parent 69f4517958
commit 11c0d63ac8
16 changed files with 398 additions and 19 deletions
+10 -2
View File
@@ -73,6 +73,7 @@ class GenerationInput:
api_key: Any
operation_type: str
prompt: str
client_device: Any | None = None
alias: str | None = None
resolution: str = "1K"
parameters: Mapping[str, Any] = field(default_factory=dict)
@@ -86,6 +87,7 @@ class GenerationInput:
class PreparedGeneration:
user: Any
api_key: Any
client_device: Any | None
operation_type: str
prompt: str
alias: str
@@ -139,11 +141,12 @@ IMAGE_URL_ALLOWED_SCHEMES = {"http", "https"}
IMAGE_URL_CHUNK_SIZE = 64 * 1024
def generate_title_response(*, user, api_key, request_data: Mapping[str, Any]) -> dict:
def generate_title_response(*, user, api_key, client_device=None, request_data: Mapping[str, Any]) -> dict:
result = run_synchronous_generation(
GenerationInput(
user=user,
api_key=api_key,
client_device=client_device,
operation_type=CallRecord.OperationType.TITLE,
prompt=request_data["prompt"],
alias=request_data.get("model") or None,
@@ -160,6 +163,7 @@ def generate_image_response(
*,
user,
api_key,
client_device=None,
request_data: Mapping[str, Any],
image_url_builder: ImageUrlBuilder | None = None,
) -> dict:
@@ -167,6 +171,7 @@ def generate_image_response(
GenerationInput(
user=user,
api_key=api_key,
client_device=client_device,
operation_type=CallRecord.OperationType.IMAGE,
prompt=request_data["prompt"],
alias=request_data.get("model") or None,
@@ -182,11 +187,12 @@ def generate_image_response(
return result.as_response_data()
def analyze_images_response(*, user, api_key, request_data: Mapping[str, Any]) -> dict:
def analyze_images_response(*, user, api_key, client_device=None, request_data: Mapping[str, Any]) -> dict:
result = run_synchronous_generation(
GenerationInput(
user=user,
api_key=api_key,
client_device=client_device,
operation_type=CallRecord.OperationType.VISION,
prompt=request_data["prompt"],
alias=request_data.get("model") or None,
@@ -257,6 +263,7 @@ def prepare_generation(generation_input: GenerationInput) -> PreparedGeneration:
return PreparedGeneration(
user=generation_input.user,
api_key=generation_input.api_key,
client_device=generation_input.client_device,
operation_type=operation_type,
prompt=prompt,
alias=model_alias.alias,
@@ -276,6 +283,7 @@ def precharge_generation(prepared: PreparedGeneration) -> PrechargedGeneration:
charge = precharge_or_raise(
user=prepared.user,
api_key=prepared.api_key,
client_device=prepared.client_device,
operation_type=prepared.operation_type,
alias=prepared.alias,
model_used=prepared.resolved_model.model,
+2
View File
@@ -38,6 +38,7 @@ def create_image_generation_task(
*,
user,
api_key,
client_device=None,
request_data: Mapping[str, Any],
idempotency_key: str = "",
) -> tuple[ImageGenerationTask, bool]:
@@ -53,6 +54,7 @@ def create_image_generation_task(
GenerationInput(
user=user,
api_key=api_key,
client_device=client_device,
operation_type=CallRecord.OperationType.IMAGE,
prompt=request_data["prompt"],
alias=request_data.get("model") or None,
+9
View File
@@ -68,12 +68,21 @@ def build_generation_route_usage_event(
) -> dict[str, Any]:
api_key = getattr(request, "auth", None)
user = getattr(request, "user", None)
client_device = getattr(request, "client_device", None)
return {
"event": EVENT_NAME,
"route_type": normalize_text(route_type, 16),
"api_key_id": getattr(api_key, "id", None),
"api_key_prefix": normalize_text(getattr(api_key, "key_prefix", ""), 32),
"user_id": getattr(user, "id", None),
"product_code": normalize_text(
getattr(client_device, "product_code", ""),
32,
),
"client_device_id": getattr(client_device, "id", None),
"device_session_present": bool(
getattr(request, "device_session_present", False)
),
"client_version": normalize_text(
request.headers.get("X-Client-Version", ""),
MAX_CLIENT_VERSION_LENGTH,
+138
View File
@@ -62,6 +62,7 @@ from apps.billing.services import RechargePayment
from apps.moderation.models import SensitiveWord
from apps.moderation.providers.keyword import reset_keyword_matcher_cache
from apps.portal.models import DownloadRelease
from apps.licensing.services import register_device
from apps.users.models import ApiKey
from apps.users.models import UserWallet
@@ -1395,6 +1396,21 @@ class GenerateApiTests(TestCase):
def auth_header(self) -> dict:
return {"HTTP_AUTHORIZATION": f"Bearer {self.raw_key}"}
def register_device_session(self, *, user=None, api_key=None):
user = user or self.user
api_key = api_key or self.api_key
result = register_device(
user=user,
api_key=api_key,
product_code="cmshopee",
device_id_version="v1",
device_id=f"test-device-{uuid.uuid4().hex}",
public_key=f"test-public-key-{uuid.uuid4().hex}",
platform="windows",
client_version="0.1.0",
)
return result.device, result.session_token
def post_with_provider(self, path, payload, provider=None, **extra):
with patch("apps.api.generation.get_provider", return_value=provider or self.provider):
return self.client.post(path, payload, format="json", **self.auth_header(), **extra)
@@ -1417,6 +1433,9 @@ class GenerateApiTests(TestCase):
"api_key_id",
"api_key_prefix",
"user_id",
"product_code",
"client_device_id",
"device_session_present",
"client_version",
"alias",
"status",
@@ -1525,6 +1544,97 @@ class GenerateApiTests(TestCase):
1,
)
def test_generation_endpoints_link_call_records_to_valid_device_session(self):
device, session_token = self.register_device_session()
device_header = {"HTTP_X_DEVICE_SESSION": session_token}
encoded = base64.b64encode(b"device-linked-image").decode("ascii")
title = self.post_with_provider(
"/api/v1/generate/title",
{"prompt": "生成标题", "model": self.title_alias},
**device_header,
)
vision = self.post_with_provider(
"/api/v1/analyze/images",
{
"prompt": "理解商品图",
"model": self.vision_alias,
"images": [{"image_base64": encoded}],
},
**device_header,
)
image = self.post_with_provider(
"/api/v1/generate/image",
{
"prompt": "生成图片",
"model": self.image_alias,
"image_base64": encoded,
},
**device_header,
)
task_submit = self.post_with_provider(
"/api/v1/generate/image/tasks",
{"prompt": "异步生成图片", "model": self.image_alias},
**device_header,
)
for response in (title, vision, image):
self.assertEqual(response.status_code, 200)
self.assertEqual(
CallRecord.objects.get(pk=response.data["call_id"]).client_device_id,
device.id,
)
self.assertEqual(task_submit.status_code, 202)
task = ImageGenerationTask.objects.get(task_id=task_submit.data["task_id"])
self.assertEqual(task.call_record.client_device_id, device.id)
poll = self.client.get(
f"/api/v1/generate/image/tasks/{task.task_id}",
**self.auth_header(),
)
self.assertEqual(poll.status_code, 200)
def test_generation_without_device_session_remains_legacy_compatible(self):
response = self.post_with_provider(
"/api/v1/generate/title",
{"prompt": "无设备头标题", "model": self.title_alias},
)
self.assertEqual(response.status_code, 200)
call = CallRecord.objects.get(pk=response.data["call_id"])
self.assertIsNone(call.client_device_id)
self.assertEqual(response.data["points_cost"], 2)
self.assertEqual(response.data["points_balance"], 98)
def test_invalid_or_cross_user_device_session_is_rejected_before_charge(self):
invalid = self.post_with_provider(
"/api/v1/generate/title",
{"prompt": "无效会话", "model": self.title_alias},
HTTP_X_DEVICE_SESSION="dvs_cmhub_invalid",
)
self.assertEqual(invalid.status_code, 401)
self.assertEqual(invalid.data["error"]["code"], "device_session_invalid")
self.assert_generation_not_charged()
other_user = get_user_model().objects.create_user(
username=f"other-device-user-{uuid.uuid4().hex[:8]}",
email=f"other-device-user-{uuid.uuid4().hex[:8]}@example.com",
password="password",
)
other_key, _raw_other_key = ApiKey.create_for_user(other_user, name="other-device")
_device, other_session_token = self.register_device_session(
user=other_user,
api_key=other_key,
)
cross_user = self.post_with_provider(
"/api/v1/generate/title",
{"prompt": "跨账号会话", "model": self.title_alias},
HTTP_X_DEVICE_SESSION=other_session_token,
)
self.assertEqual(cross_user.status_code, 403)
self.assertEqual(cross_user.data["error"]["code"], "device_mismatch")
self.assert_generation_not_charged()
def test_analyze_images_supports_single_image_with_explicit_alias(self):
encoded = base64.b64encode(b"single-image").decode("ascii")
@@ -1963,6 +2073,9 @@ class GenerateApiTests(TestCase):
self.assertEqual(event["api_key_id"], self.api_key.id)
self.assertEqual(event["api_key_prefix"], self.api_key.key_prefix)
self.assertEqual(event["user_id"], self.user.id)
self.assertEqual(event["product_code"], "")
self.assertIsNone(event["client_device_id"])
self.assertFalse(event["device_session_present"])
self.assertEqual(event["client_version"], "0.1.1")
self.assertEqual(event["alias"], self.image_alias)
self.assertEqual(event["status"], "success")
@@ -1972,6 +2085,31 @@ class GenerateApiTests(TestCase):
self.assertGreaterEqual(event["latency_ms"], 0)
self.assert_generation_telemetry_is_safe(event, payload=payload)
def test_sync_image_usage_telemetry_records_only_safe_device_metadata(self):
device, session_token = self.register_device_session()
encoded = base64.b64encode(b"telemetry-device-image").decode("ascii")
payload = {
"prompt": "设备遥测图片生成",
"model": self.image_alias,
"image_base64": encoded,
}
with self.assertLogs("cmhub.api.generation_usage", level="INFO") as captured:
response = self.post_with_provider(
"/api/v1/generate/image",
payload,
HTTP_X_DEVICE_SESSION=session_token,
HTTP_X_CLIENT_VERSION="0.1.3",
)
self.assertEqual(response.status_code, 200)
event = self.telemetry_event_from_logs(captured)
self.assertEqual(event["product_code"], "cmshopee")
self.assertEqual(event["client_device_id"], device.id)
self.assertTrue(event["device_session_present"])
self.assertNotIn(session_token, json.dumps(event, ensure_ascii=False))
self.assert_generation_telemetry_is_safe(event, payload=payload)
def test_async_image_submit_usage_telemetry_logs_safe_success_event(self):
payload = {
"prompt": "生成异步图片遥测测试",
+132
View File
@@ -67,8 +67,10 @@ from apps.portal.models import DownloadRelease
from apps.licensing.authentication import DeviceSessionAuthentication
from apps.licensing.services import (
DeviceRegistrationError,
DeviceSessionValidationError,
record_device_heartbeat,
register_device,
resolve_optional_device_session,
)
logger = logging.getLogger(__name__)
@@ -84,13 +86,55 @@ class ExternalApiView(APIView):
raise AuthenticationFailed(api_error("unauthorized", "缺失或无效 API Key"))
super().permission_denied(request, message=message, code=code)
def optional_client_device(self, request):
raw_token = request.headers.get("X-Device-Session", "")
request.device_session_present = bool(raw_token.strip())
try:
session = resolve_optional_device_session(
user=request.user,
raw_token=raw_token,
)
except DeviceSessionValidationError as exc:
http_status = (
status.HTTP_401_UNAUTHORIZED
if exc.code == "device_session_invalid"
else status.HTTP_403_FORBIDDEN
)
raise ApiRequestError(exc.code, exc.message, http_status) from exc
request.client_device = session.device if session is not None else None
return request.client_device
class GenerateTitleView(ExternalApiView):
throttle_classes = (GenerateRateThrottle,)
def post(self, request):
started = telemetry_start_time()
alias = request_alias(request.data)
try:
client_device = self.optional_client_device(request)
except ApiRequestError as exc:
log_generation_route_usage(
route_type="title",
request=request,
alias=alias,
status="error",
latency_ms=telemetry_elapsed_ms(started),
error_code=exc.code,
http_status=exc.http_status,
)
return Response(exc.as_response_data(), status=exc.http_status)
serializer = GenerateTitleRequestSerializer(data=request.data)
if not serializer.is_valid():
log_generation_route_usage(
route_type="title",
request=request,
alias=alias,
status="error",
latency_ms=telemetry_elapsed_ms(started),
error_code="bad_request",
http_status=status.HTTP_400_BAD_REQUEST,
)
return Response(
api_error("bad_request", "参数错误"),
status=status.HTTP_400_BAD_REQUEST,
@@ -99,10 +143,28 @@ class GenerateTitleView(ExternalApiView):
data = generate_title_response(
user=request.user,
api_key=request.auth,
client_device=client_device,
request_data=serializer.validated_data,
)
except ApiRequestError as exc:
log_generation_route_usage(
route_type="title",
request=request,
alias=alias,
status="error",
latency_ms=telemetry_elapsed_ms(started),
error_code=exc.code,
http_status=exc.http_status,
)
return Response(exc.as_response_data(), status=exc.http_status)
log_generation_route_usage(
route_type="title",
request=request,
alias=data.get("alias") or alias,
status="success",
latency_ms=telemetry_elapsed_ms(started),
http_status=status.HTTP_200_OK,
)
return Response(data, status=status.HTTP_200_OK)
@@ -110,8 +172,32 @@ class AnalyzeImagesView(ExternalApiView):
throttle_classes = (GenerateRateThrottle,)
def post(self, request):
started = telemetry_start_time()
alias = request_alias(request.data)
try:
client_device = self.optional_client_device(request)
except ApiRequestError as exc:
log_generation_route_usage(
route_type="vision",
request=request,
alias=alias,
status="error",
latency_ms=telemetry_elapsed_ms(started),
error_code=exc.code,
http_status=exc.http_status,
)
return Response(exc.as_response_data(), status=exc.http_status)
serializer = AnalyzeImagesRequestSerializer(data=request.data)
if not serializer.is_valid():
log_generation_route_usage(
route_type="vision",
request=request,
alias=alias,
status="error",
latency_ms=telemetry_elapsed_ms(started),
error_code="bad_request",
http_status=status.HTTP_400_BAD_REQUEST,
)
return Response(
api_error("bad_request", "参数错误"),
status=status.HTTP_400_BAD_REQUEST,
@@ -120,10 +206,28 @@ class AnalyzeImagesView(ExternalApiView):
data = analyze_images_response(
user=request.user,
api_key=request.auth,
client_device=client_device,
request_data=serializer.validated_data,
)
except ApiRequestError as exc:
log_generation_route_usage(
route_type="vision",
request=request,
alias=alias,
status="error",
latency_ms=telemetry_elapsed_ms(started),
error_code=exc.code,
http_status=exc.http_status,
)
return Response(exc.as_response_data(), status=exc.http_status)
log_generation_route_usage(
route_type="vision",
request=request,
alias=data.get("alias") or alias,
status="success",
latency_ms=telemetry_elapsed_ms(started),
http_status=status.HTTP_200_OK,
)
return Response(data, status=status.HTTP_200_OK)
@@ -133,6 +237,19 @@ class GenerateImageView(ExternalApiView):
def post(self, request):
started = telemetry_start_time()
alias = request_alias(request.data)
try:
client_device = self.optional_client_device(request)
except ApiRequestError as exc:
log_generation_route_usage(
route_type="sync",
request=request,
alias=alias,
status="error",
latency_ms=telemetry_elapsed_ms(started),
error_code=exc.code,
http_status=exc.http_status,
)
return Response(exc.as_response_data(), status=exc.http_status)
serializer = GenerateImageRequestSerializer(data=request.data)
if not serializer.is_valid():
log_generation_route_usage(
@@ -153,6 +270,7 @@ class GenerateImageView(ExternalApiView):
data = generate_image_response(
user=request.user,
api_key=request.auth,
client_device=client_device,
request_data=serializer.validated_data,
image_url_builder=request.build_absolute_uri,
)
@@ -184,6 +302,19 @@ class GenerateImageTaskSubmitView(ExternalApiView):
def post(self, request):
started = telemetry_start_time()
alias = request_alias(request.data)
try:
client_device = self.optional_client_device(request)
except ApiRequestError as exc:
log_generation_route_usage(
route_type="async",
request=request,
alias=alias,
status="error",
latency_ms=telemetry_elapsed_ms(started),
error_code=exc.code,
http_status=exc.http_status,
)
return Response(exc.as_response_data(), status=exc.http_status)
serializer = GenerateImageRequestSerializer(data=request.data)
if not serializer.is_valid():
log_generation_route_usage(
@@ -204,6 +335,7 @@ class GenerateImageTaskSubmitView(ExternalApiView):
task, _created = create_image_generation_task(
user=request.user,
api_key=request.auth,
client_device=client_device,
request_data=serializer.validated_data,
idempotency_key=request.headers.get("Idempotency-Key", ""),
)
+3 -2
View File
@@ -109,6 +109,7 @@ class CallRecordAdmin(ReadOnlyLedgerAdmin):
"created_at",
"user",
"api_key",
"client_device",
"operation_type",
"alias",
"model_used",
@@ -116,7 +117,7 @@ class CallRecordAdmin(ReadOnlyLedgerAdmin):
"status",
"upstream_latency_ms",
)
list_filter = ("operation_type", "status", "created_at")
list_filter = ("operation_type", "status", "client_device", "created_at")
search_fields = (
"user__username",
"user__email",
@@ -129,7 +130,7 @@ class CallRecordAdmin(ReadOnlyLedgerAdmin):
"result_summary",
)
ordering = ("-created_at", "-id")
list_select_related = ("user", "api_key")
list_select_related = ("user", "api_key", "client_device")
@admin.register(SignupBonusGrant)
@@ -0,0 +1,27 @@
# Generated by Django 5.2.15 on 2026-07-21 01:34
import django.db.models.deletion
from django.conf import settings
from django.db import migrations, models
class Migration(migrations.Migration):
dependencies = [
('billing', '0009_signup_bonus_default_ten'),
('licensing', '0001_initial'),
('users', '0005_alter_apikey_created_at_alter_apikey_key_hash_and_more'),
migrations.swappable_dependency(settings.AUTH_USER_MODEL),
]
operations = [
migrations.AddField(
model_name='callrecord',
name='client_device',
field=models.ForeignKey(blank=True, null=True, on_delete=django.db.models.deletion.SET_NULL, related_name='call_records', to='licensing.clientdevice', verbose_name='客户端设备'),
),
migrations.AddIndex(
model_name='callrecord',
index=models.Index(fields=['client_device', 'created_at'], name='call_record_client__d5a0fb_idx'),
),
]
+9
View File
@@ -35,6 +35,14 @@ class CallRecord(models.Model):
on_delete=models.SET_NULL,
related_name="call_records",
)
client_device = models.ForeignKey(
"licensing.ClientDevice",
verbose_name="客户端设备",
null=True,
blank=True,
on_delete=models.SET_NULL,
related_name="call_records",
)
operation_type = models.CharField("操作类型", max_length=32, choices=OperationType.choices)
alias = models.CharField("能力别名", max_length=64, blank=True)
model_used = models.CharField("实际模型", max_length=128, blank=True)
@@ -68,6 +76,7 @@ class CallRecord(models.Model):
indexes = [
models.Index(fields=("user", "created_at")),
models.Index(fields=("api_key", "created_at")),
models.Index(fields=("client_device", "created_at")),
models.Index(fields=("operation_type", "status")),
models.Index(fields=("alias",)),
]
+2
View File
@@ -504,6 +504,7 @@ def precharge_call(
model_used: str = "",
resolution: str | None = None,
api_key=None,
client_device=None,
prompt: str = "",
) -> CallCharge:
points_cost = _validate_positive_points(points_cost)
@@ -524,6 +525,7 @@ def precharge_call(
call_record = CallRecord.objects.create(
user=user,
api_key=api_key,
client_device=client_device,
operation_type=operation_type,
alias=normalized_alias,
model_used=str(model_used or "").strip(),
+28
View File
@@ -17,6 +17,10 @@ class DeviceRegistrationError(Exception):
super().__init__(message)
class DeviceSessionValidationError(DeviceRegistrationError):
pass
@dataclass(frozen=True)
class DeviceRegistrationResult:
device: ClientDevice
@@ -150,3 +154,27 @@ def record_device_heartbeat(session: DeviceSession, *, now=None) -> bool:
client_version=session.device.client_version,
)
return True
def resolve_optional_device_session(*, user, raw_token: str) -> DeviceSession | None:
raw_token = str(raw_token or "").strip()
if not raw_token:
return None
try:
session = DeviceSession.objects.select_related("device", "device__user").get(
token_hash=DeviceSession.hash_token(raw_token)
)
except DeviceSession.DoesNotExist as exc:
raise DeviceSessionValidationError(
"device_session_invalid",
"设备会话无效或已过期",
) from exc
if not session.matches_token(raw_token) or not session.is_active_at():
raise DeviceSessionValidationError("device_session_invalid", "设备会话无效或已过期")
if session.device.status != ClientDevice.Status.ACTIVE:
raise DeviceSessionValidationError("device_revoked", "设备已被吊销")
if session.device.user_id != user.id:
raise DeviceSessionValidationError("device_mismatch", "设备会话不属于当前账号")
return session