feat: associate generation calls with devices
This commit is contained in:
+10
-2
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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": "生成异步图片遥测测试",
|
||||
|
||||
@@ -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", ""),
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user