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