Compare commits
4 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| cfb5c7ab89 | |||
| d6d100baeb | |||
| 4d99c20f6b | |||
| e9c119aea4 |
@@ -29,7 +29,7 @@ import tempfile
|
||||
import threading
|
||||
import time
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
from typing import Any, Callable, Optional
|
||||
|
||||
from celery import Task, shared_task
|
||||
from celery.exceptions import Retry
|
||||
@@ -791,8 +791,12 @@ def _script_from_xml(raw: str, job: ViralVideoJob) -> dict | None:
|
||||
return base
|
||||
|
||||
|
||||
def _step_script_generation(job: ViralVideoJob, image_analysis: dict) -> dict:
|
||||
"""步骤: 意图理解 + 分镜生成一次完成(v3,模板 + XML 解析)。"""
|
||||
def _step_script_generation(
|
||||
job: ViralVideoJob,
|
||||
image_analysis: dict,
|
||||
on_delta: Optional[Callable[[str, str], None]] = None,
|
||||
) -> dict:
|
||||
"""步骤: 意图理解 + 分镜生成一次完成(v3,模板 + XML 解析,支持流式推送)。"""
|
||||
try:
|
||||
from packages.application.viral_video.prompt_loader import (
|
||||
get_template,
|
||||
@@ -928,6 +932,140 @@ def _step_script_generation(job: ViralVideoJob, image_analysis: dict) -> dict:
|
||||
|
||||
return normalized
|
||||
|
||||
# ── 流式内部函数 ──────────────────────────────────────────────────────
|
||||
def _stream_chat_with_fallback(client, messages, temp, max_tok, tmo):
|
||||
"""流式调用,失败/超时则降级到同步 chat_completion;yield 每段 delta 文本。
|
||||
|
||||
返回 (full_text, used_stream)。
|
||||
"""
|
||||
full_parts = []
|
||||
stream_ok = False
|
||||
if client and client.is_available and on_delta is not None and hasattr(client, "chat_completion_stream"):
|
||||
try:
|
||||
for chunk in client.chat_completion_stream(
|
||||
messages, temperature=temp, max_tokens=max_tok, timeout=tmo or 120
|
||||
):
|
||||
if chunk:
|
||||
full_parts.append(chunk)
|
||||
yield ("delta", chunk)
|
||||
stream_ok = True
|
||||
except Exception as e:
|
||||
logger.warning("[爆款视频] 流式调用失败,降级同步: %s", e, exc_info=True)
|
||||
full_parts = [] # 重置,走同步
|
||||
if not stream_ok:
|
||||
raw = client.chat_completion(messages, temperature=temp, max_tokens=max_tok, timeout=tmo)
|
||||
if raw:
|
||||
full_parts = [raw]
|
||||
yield ("delta", raw) # 一次性推送完整文本(同步回退)
|
||||
else:
|
||||
full_parts = []
|
||||
yield ("done", "".join(full_parts))
|
||||
|
||||
def _try_gen_stream(client, temp: float, max_tok: int, label: str, tmo: int, user_text: str = None):
|
||||
"""流式版本 _try_gen:边收 token 边调 on_delta,最终返回 normalized dict 或 None。"""
|
||||
if not client or not client.is_available:
|
||||
return None
|
||||
_u = user_text if user_text is not None else user
|
||||
messages = [{"role": "system", "content": system}, {"role": "user", "content": _u}]
|
||||
logger.info("[爆款视频] 分镜生成(流式) model=%s label=%s timeout=%d", client.model, label, tmo)
|
||||
|
||||
full_text_buf = []
|
||||
pending_delta_buf = []
|
||||
last_emit = 0.0
|
||||
MIN_INTERVAL = 0.25 # 至少 250ms 一次,约 4 次/秒
|
||||
MIN_CHARS = 40 # 累积 ~40 字符才推送
|
||||
|
||||
def _flush(force: bool = False):
|
||||
nonlocal last_emit, pending_delta_buf
|
||||
if not pending_delta_buf:
|
||||
return
|
||||
now = time.time()
|
||||
if not force and (now - last_emit) < MIN_INTERVAL:
|
||||
return
|
||||
delta_text = "".join(pending_delta_buf)
|
||||
pending_delta_buf = []
|
||||
full_text = "".join(full_text_buf)
|
||||
last_emit = now
|
||||
try:
|
||||
on_delta(delta_text, full_text)
|
||||
except Exception as e:
|
||||
logger.warning("[爆款视频] on_delta 回调失败: %s", e)
|
||||
|
||||
final_raw = None
|
||||
try:
|
||||
for kind, payload in _stream_chat_with_fallback(client, messages, temp, max_tok, tmo):
|
||||
if kind == "delta":
|
||||
full_text_buf.append(payload)
|
||||
pending_delta_buf.append(payload)
|
||||
# 判断是否触发推送
|
||||
buf_text = "".join(pending_delta_buf)
|
||||
should_flush = False
|
||||
if len(buf_text) >= MIN_CHARS:
|
||||
should_flush = True
|
||||
elif any(tok in buf_text for tok in ("\n", "</", "/>", ">\n")):
|
||||
# 换行或 XML 标签闭合时尽早 flush
|
||||
if len(buf_text) >= 10:
|
||||
should_flush = True
|
||||
if should_flush:
|
||||
_flush(force=False)
|
||||
elif kind == "done":
|
||||
final_raw = payload
|
||||
except Exception as e:
|
||||
logger.warning("[爆款视频] 流式生成异常 label=%s err=%s", label, e, exc_info=True)
|
||||
return None
|
||||
|
||||
_flush(force=True) # 剩余全部推送
|
||||
|
||||
if not final_raw:
|
||||
return None
|
||||
|
||||
normalized = _script_from_xml(final_raw, job)
|
||||
if normalized is None:
|
||||
parsed_json = _safe_json_loads(final_raw)
|
||||
if isinstance(parsed_json, dict):
|
||||
normalized = _validate_and_normalize_script(parsed_json, job)
|
||||
else:
|
||||
return None
|
||||
voiceover = normalized.get("voiceover_script") or ""
|
||||
shots_cnt = len(normalized.get("shots") or [])
|
||||
is_fallback = shots_cnt < 1 or len(voiceover) < 12
|
||||
logger.info(
|
||||
"[爆款视频] 分镜结果(流式) label=%s voiceover_len=%d shots_cnt=%d fallback=%s",
|
||||
label,
|
||||
len(voiceover),
|
||||
shots_cnt,
|
||||
is_fallback,
|
||||
)
|
||||
if is_fallback:
|
||||
return None
|
||||
_dur = max(5, min(30, int(getattr(job, "duration", 15) or 15)))
|
||||
_max_chars = _dur * 3
|
||||
_voiceover_chars = len(voiceover.strip())
|
||||
_min_chars = max(10, int(_dur * 2.2))
|
||||
if _voiceover_chars > _max_chars:
|
||||
logger.warning(
|
||||
"[爆款视频] 口播超长(流式) label=%s voiceover_chars=%d max=%d", label, _voiceover_chars, _max_chars
|
||||
)
|
||||
return None
|
||||
if _voiceover_chars < _min_chars:
|
||||
logger.warning(
|
||||
"[爆款视频] 口播过短(流式) label=%s voiceover_chars=%d min=%d", label, _voiceover_chars, _min_chars
|
||||
)
|
||||
return None
|
||||
_shots = normalized.get("shots") or []
|
||||
_shot_count = len(_shots)
|
||||
_expected_range = _get_expected_shot_count(_dur)
|
||||
if _expected_range and (_shot_count < _expected_range[0] or _shot_count > _expected_range[1]):
|
||||
logger.warning(
|
||||
"[爆款视频] 镜头数量不符(流式) label=%s shots=%d expected=%s", label, _shot_count, _expected_range
|
||||
)
|
||||
return None
|
||||
_time_valid = _validate_shot_timeline(_shots, _dur)
|
||||
if not _time_valid:
|
||||
logger.warning("[爆款视频] 时间轴不合法(流式) label=%s dur=%ds", label, _dur)
|
||||
return None
|
||||
return normalized
|
||||
|
||||
client_fast = ai_router.get_llm_client("storyboard", variant="primary")
|
||||
client_pro = ai_router.get_llm_client("storyboard", variant="lite")
|
||||
fast_tmo = int(os.environ.get("VIRAL_VIDEO_SCRIPT_FAST_TIMEOUT", "90"))
|
||||
@@ -940,13 +1078,16 @@ def _step_script_generation(job: ViralVideoJob, image_analysis: dict) -> dict:
|
||||
_char_hint = f"口播总字数严格控制在 {_min_chars}~{_max_chars} 字({_dur}秒视频),超长会导致配音失败"
|
||||
user = user + "\n\n" + _char_hint
|
||||
|
||||
# 选择 _try_gen 实现:有 on_delta 用流式,否则保持原同步逻辑
|
||||
_do_gen = _try_gen_stream if on_delta is not None else _try_gen
|
||||
|
||||
try:
|
||||
result = _try_gen(client_fast, 0.8, 2500, "fast-first", fast_tmo)
|
||||
result = _do_gen(client_fast, 0.8, 2500, "fast-first", fast_tmo)
|
||||
if result is not None:
|
||||
return result
|
||||
if time.time() > deadline:
|
||||
return _finalize_fallback_script(job)
|
||||
result = _try_gen(client_fast, 0.6, 3200, "fast-retry", fast_tmo)
|
||||
result = _do_gen(client_fast, 0.6, 3200, "fast-retry", fast_tmo)
|
||||
if result is not None:
|
||||
return result
|
||||
# Bug2: 压缩重试 — 用更严格约束要求 LLM 压缩口播
|
||||
@@ -954,12 +1095,12 @@ def _step_script_generation(job: ViralVideoJob, image_analysis: dict) -> dict:
|
||||
user_compressed = user + "\n\n【紧急】上一次生成口播超长,请将口播压缩到 {} 字以内,保留核心卖点。".format(
|
||||
_max_chars
|
||||
)
|
||||
compressed_result = _try_gen(client_fast, 0.5, 2000, "compress-retry", fast_tmo, user_text=user_compressed)
|
||||
compressed_result = _do_gen(client_fast, 0.5, 2000, "compress-retry", fast_tmo, user_text=user_compressed)
|
||||
if compressed_result is not None:
|
||||
return compressed_result
|
||||
if client_pro and client_pro.is_available and client_pro.model != client_fast.model:
|
||||
if time.time() <= deadline:
|
||||
result = _try_gen(client_pro, 0.7, 3500, "pro-fallback", pro_tmo)
|
||||
result = _do_gen(client_pro, 0.7, 3500, "pro-fallback", pro_tmo)
|
||||
if result is not None:
|
||||
return result
|
||||
return _finalize_fallback_script(job)
|
||||
@@ -2046,10 +2187,35 @@ def run_viral_video_generate_copy(self: Task, job_id: str) -> dict:
|
||||
_save_job(repo, job, session)
|
||||
_hb_stop, _hb_thread = _start_heartbeat_thread(job_id)
|
||||
|
||||
# v8:意图理解并入分镜生成,一次 LLM 调用
|
||||
# v8:意图理解并入分镜生成,一次 LLM 调用;流式推送 script_delta 给前端
|
||||
image_analysis = normalize_image_analysis(job.image_analysis)
|
||||
_set_stage(job, repo, session, ViralVideoStage.SCRIPT_GENERATION, "正在编排分镜脚本...")
|
||||
copy_result = _step_script_generation(job, image_analysis)
|
||||
|
||||
_script_full_text = []
|
||||
_script_last_emit_ts = [0.0]
|
||||
_script_stage_start = time.time()
|
||||
|
||||
def _on_script_delta(delta: str, full_text: str):
|
||||
"""流式回调:推 viral_video:script_delta 事件到 Redis pub/sub(WS 桥接前端)。"""
|
||||
_script_full_text.append(delta) if delta else None
|
||||
now = time.time()
|
||||
# 速率保护:on_delta 已经做了基础节流;这里再加一道 200ms 兜底防止消息风暴
|
||||
if now - _script_last_emit_ts[0] < 0.2:
|
||||
return
|
||||
_script_last_emit_ts[0] = now
|
||||
# 进度估算:基于 full_text 长度线性增长(最长 ~3500 字 = 70% 进度位)
|
||||
_est_progress = min(70.0, 15.0 + (len(full_text) / 3500.0) * 55.0)
|
||||
_elapsed = now - _script_stage_start
|
||||
_emit_progress(
|
||||
job_id,
|
||||
ViralVideoStage.SCRIPT_GENERATION,
|
||||
round(_est_progress, 1),
|
||||
f"正在编排分镜脚本...({len(full_text)}字,{_elapsed:.0f}s)",
|
||||
{"delta": delta, "full_text": full_text, "text_length": len(full_text)},
|
||||
event_type="viral_video:script_delta",
|
||||
)
|
||||
|
||||
copy_result = _step_script_generation(job, image_analysis, on_delta=_on_script_delta)
|
||||
_emit_progress(
|
||||
job_id,
|
||||
ViralVideoStage.SCRIPT_GENERATION,
|
||||
|
||||
@@ -363,6 +363,111 @@ class DoubaoClient:
|
||||
logger.error("豆包API调用最终失败: elapsed=%.1fs err=%s", time.time() - _t0, last_error)
|
||||
return None
|
||||
|
||||
def chat_completion_stream(
|
||||
self,
|
||||
messages: list[dict[str, str]],
|
||||
temperature: float = 0.7,
|
||||
max_tokens: int | None = None,
|
||||
model: str | None = None,
|
||||
timeout: int | None = None,
|
||||
**kwargs,
|
||||
):
|
||||
"""流式调用 Chat Completion 接口(SSE),逐块 yield delta 文本。
|
||||
|
||||
Yields:
|
||||
str: 增量文本片段(delta.content);全部结束后 StopIteration。
|
||||
失败时 yield 空并返回(由调用方决定是否降级到同步调用)。
|
||||
|
||||
注意:
|
||||
- 流式不做 finish_reason=length 自动扩容(流式难以拼接重试);
|
||||
如果调用方需要 length 截断处理,建议自行 fallback 到同步 chat_completion。
|
||||
- 重试只在连接建立阶段(首包之前)有效;一旦开始 yield,错误直接抛出。
|
||||
"""
|
||||
if not self.is_available:
|
||||
return
|
||||
|
||||
import json as _json
|
||||
|
||||
url = f"{self.base_url}/chat/completions"
|
||||
headers = {
|
||||
"Authorization": f"Bearer {self.api_key}",
|
||||
"Content-Type": "application/json",
|
||||
"Accept": "text/event-stream",
|
||||
}
|
||||
effective_max_tokens = max_tokens if max_tokens is not None else (self.max_tokens or 1024)
|
||||
payload: dict[str, Any] = {
|
||||
"model": model or self.model,
|
||||
"messages": messages,
|
||||
"temperature": temperature,
|
||||
"max_tokens": effective_max_tokens,
|
||||
"stream": True,
|
||||
}
|
||||
if self.extra_params:
|
||||
payload.update(self.extra_params)
|
||||
if kwargs:
|
||||
payload.update(kwargs)
|
||||
|
||||
_t0 = time.time()
|
||||
for attempt in range(self.max_retries + 1):
|
||||
try:
|
||||
_req_timeout = httpx.Timeout(connect=10.0, read=120.0, write=30.0, pool=10.0)
|
||||
if timeout:
|
||||
_req_timeout = httpx.Timeout(connect=10.0, read=max(int(timeout), 30), write=30.0, pool=10.0)
|
||||
with httpx.stream(
|
||||
"POST",
|
||||
url,
|
||||
headers=headers,
|
||||
json=payload,
|
||||
timeout=_req_timeout,
|
||||
) as resp:
|
||||
resp.raise_for_status()
|
||||
for line in resp.iter_lines():
|
||||
if not line:
|
||||
continue
|
||||
line = line.strip()
|
||||
if not line.startswith("data:"):
|
||||
continue
|
||||
data_str = line[5:].strip()
|
||||
if data_str == "[DONE]":
|
||||
break
|
||||
try:
|
||||
chunk = _json.loads(data_str)
|
||||
except (_json.JSONDecodeError, ValueError):
|
||||
continue
|
||||
choices = chunk.get("choices") or []
|
||||
if not choices:
|
||||
continue
|
||||
delta = choices[0].get("delta") or {}
|
||||
content_piece = delta.get("content") or ""
|
||||
if content_piece:
|
||||
yield content_piece
|
||||
finish_reason = choices[0].get("finish_reason")
|
||||
if finish_reason:
|
||||
self.last_finish_reason = finish_reason
|
||||
_elapsed = time.time() - _t0
|
||||
logger.info(
|
||||
"[doubao] chat_completion_stream 完成 model=%s elapsed=%.1fs attempt=%d",
|
||||
payload.get("model"),
|
||||
_elapsed,
|
||||
attempt + 1,
|
||||
)
|
||||
return
|
||||
except Exception as e:
|
||||
if attempt < self.max_retries:
|
||||
wait = 0.5 * (2**attempt)
|
||||
logger.warning(
|
||||
"豆包流式API失败,%.1fs后重试 (%d/%d, elapsed=%.1fs): %s",
|
||||
wait,
|
||||
attempt + 1,
|
||||
self.max_retries + 1,
|
||||
time.time() - _t0,
|
||||
e,
|
||||
)
|
||||
time.sleep(wait)
|
||||
continue
|
||||
logger.error("豆包流式API最终失败: elapsed=%.1fs err=%s", time.time() - _t0, e)
|
||||
return
|
||||
|
||||
def vision_completion(
|
||||
self,
|
||||
messages: list[dict],
|
||||
|
||||
@@ -196,3 +196,159 @@ class TestGetDoubaoClient:
|
||||
"""返回 DoubaoClient 实例"""
|
||||
client = get_doubao_client()
|
||||
assert isinstance(client, DoubaoClient)
|
||||
|
||||
|
||||
class TestChatCompletionStream:
|
||||
"""chat_completion_stream 方法测试"""
|
||||
|
||||
def _fake_sse(self, chunks: list[str]):
|
||||
"""构造一个伪装的 httpx stream 响应,按 chunks 逐行返回 SSE。"""
|
||||
import json as _json
|
||||
|
||||
lines = []
|
||||
for piece in chunks:
|
||||
evt = {"choices": [{"delta": {"content": piece}, "finish_reason": None}]}
|
||||
lines.append("data: " + _json.dumps(evt, ensure_ascii=False))
|
||||
lines.append("data: " + _json.dumps({"choices": [{"delta": {}, "finish_reason": "stop"}]}))
|
||||
lines.append("data: [DONE]")
|
||||
|
||||
m = MagicMock()
|
||||
m.status_code = 200
|
||||
m.__enter__ = MagicMock(return_value=m)
|
||||
m.__exit__ = MagicMock(return_value=False)
|
||||
m.raise_for_status = MagicMock()
|
||||
m.iter_lines = MagicMock(return_value=iter(lines))
|
||||
return m
|
||||
|
||||
def test_stream_yields_delta_content(self, client_with_key):
|
||||
"""流式调用逐段 yield delta.content"""
|
||||
import json as _json
|
||||
|
||||
fake = self._fake_sse(["你好", ",", "世界"])
|
||||
with patch("packages.shared.ai_client.httpx.stream", return_value=fake) as mock_stream:
|
||||
out = list(
|
||||
client_with_key.chat_completion_stream(
|
||||
messages=[{"role": "user", "content": "hi"}], max_tokens=100, timeout=30
|
||||
)
|
||||
)
|
||||
assert out == ["你好", ",", "世界"]
|
||||
call_kwargs = mock_stream.call_args[1]
|
||||
assert call_kwargs["json"]["stream"] is True
|
||||
assert call_kwargs["json"]["max_tokens"] == 100
|
||||
|
||||
def test_stream_unavailable_returns_empty(self, client_without_key):
|
||||
"""不可用时返回空生成器(不调用 httpx.stream)"""
|
||||
with patch("packages.shared.ai_client.httpx.stream") as mock_stream:
|
||||
out = list(client_without_key.chat_completion_stream(messages=[{"role": "user", "content": "hi"}]))
|
||||
assert out == []
|
||||
mock_stream.assert_not_called()
|
||||
|
||||
def test_stream_handles_done_marker(self, client_with_key):
|
||||
"""遇到 [DONE] 正确终止,不把它当内容"""
|
||||
import json as _json
|
||||
|
||||
lines = [
|
||||
"data: " + _json.dumps({"choices": [{"delta": {"content": "A"}}]}),
|
||||
"data: [DONE]",
|
||||
"data: " + _json.dumps({"choices": [{"delta": {"content": "NEVER"}}]}), # 应被忽略
|
||||
]
|
||||
m = MagicMock()
|
||||
m.status_code = 200
|
||||
m.__enter__ = MagicMock(return_value=m)
|
||||
m.__exit__ = MagicMock(return_value=False)
|
||||
m.raise_for_status = MagicMock()
|
||||
m.iter_lines = MagicMock(return_value=iter(lines))
|
||||
with patch("packages.shared.ai_client.httpx.stream", return_value=m):
|
||||
out = list(
|
||||
client_with_key.chat_completion_stream(
|
||||
messages=[{"role": "user", "content": "hi"}], max_tokens=50, timeout=30
|
||||
)
|
||||
)
|
||||
assert out == ["A"]
|
||||
|
||||
def test_stream_skips_empty_delta(self, client_with_key):
|
||||
"""空 delta(role 等 metadata)不应产出内容"""
|
||||
import json as _json
|
||||
|
||||
lines = [
|
||||
"data: " + _json.dumps({"choices": [{"delta": {"role": "assistant"}}]}),
|
||||
"data: " + _json.dumps({"choices": [{"delta": {"content": "hi"}}]}),
|
||||
"data: " + _json.dumps({"choices": [{"delta": {}}]}),
|
||||
"data: [DONE]",
|
||||
]
|
||||
m = MagicMock()
|
||||
m.status_code = 200
|
||||
m.__enter__ = MagicMock(return_value=m)
|
||||
m.__exit__ = MagicMock(return_value=False)
|
||||
m.raise_for_status = MagicMock()
|
||||
m.iter_lines = MagicMock(return_value=iter(lines))
|
||||
with patch("packages.shared.ai_client.httpx.stream", return_value=m):
|
||||
out = list(
|
||||
client_with_key.chat_completion_stream(
|
||||
messages=[{"role": "user", "content": "x"}], max_tokens=50, timeout=30
|
||||
)
|
||||
)
|
||||
assert out == ["hi"]
|
||||
|
||||
def test_stream_invalid_json_lines_skipped(self, client_with_key):
|
||||
"""SSE 行里脏数据/非 JSON 不应中断流"""
|
||||
import json as _json
|
||||
|
||||
lines = [
|
||||
"data: " + _json.dumps({"choices": [{"delta": {"content": "ok"}}]}),
|
||||
"data: not-a-json",
|
||||
":comment line",
|
||||
"",
|
||||
"event: ping",
|
||||
"data: " + _json.dumps({"choices": [{"delta": {"content": "2"}}]}),
|
||||
"data: [DONE]",
|
||||
]
|
||||
m = MagicMock()
|
||||
m.status_code = 200
|
||||
m.__enter__ = MagicMock(return_value=m)
|
||||
m.__exit__ = MagicMock(return_value=False)
|
||||
m.raise_for_status = MagicMock()
|
||||
m.iter_lines = MagicMock(return_value=iter(lines))
|
||||
with patch("packages.shared.ai_client.httpx.stream", return_value=m):
|
||||
out = list(
|
||||
client_with_key.chat_completion_stream(
|
||||
messages=[{"role": "user", "content": "x"}], max_tokens=50, timeout=30
|
||||
)
|
||||
)
|
||||
assert out == ["ok", "2"]
|
||||
|
||||
def test_stream_network_error_yields_empty(self, client_with_key):
|
||||
"""网络错误(重试耗尽)yield 空,不抛异常给调用方"""
|
||||
import httpx as _httpx
|
||||
|
||||
# max_retries=2 → 3 次总尝试
|
||||
client = client_with_key
|
||||
client.max_retries = 1 # 只重试 1 次,缩短测试
|
||||
with patch("packages.shared.ai_client.httpx.stream", side_effect=_httpx.ConnectError("boom")):
|
||||
out = list(
|
||||
client.chat_completion_stream(messages=[{"role": "user", "content": "x"}], max_tokens=50, timeout=5)
|
||||
)
|
||||
assert out == []
|
||||
|
||||
def test_stream_records_finish_reason(self, client_with_key):
|
||||
"""流式结束后 last_finish_reason 被正确记录"""
|
||||
import json as _json
|
||||
|
||||
lines = [
|
||||
"data: " + _json.dumps({"choices": [{"delta": {"content": "x"}, "finish_reason": None}]}),
|
||||
"data: " + _json.dumps({"choices": [{"delta": {}, "finish_reason": "stop"}]}),
|
||||
"data: [DONE]",
|
||||
]
|
||||
m = MagicMock()
|
||||
m.status_code = 200
|
||||
m.__enter__ = MagicMock(return_value=m)
|
||||
m.__exit__ = MagicMock(return_value=False)
|
||||
m.raise_for_status = MagicMock()
|
||||
m.iter_lines = MagicMock(return_value=iter(lines))
|
||||
with patch("packages.shared.ai_client.httpx.stream", return_value=m):
|
||||
list(
|
||||
client_with_key.chat_completion_stream(
|
||||
messages=[{"role": "user", "content": "x"}], max_tokens=50, timeout=30
|
||||
)
|
||||
)
|
||||
assert client_with_key.last_finish_reason == "stop"
|
||||
|
||||
@@ -483,3 +483,88 @@ class TestWSInitialSnapshot:
|
||||
received, _, _ = _run_ws_handshake(job=job)
|
||||
data = received[0]["data"]
|
||||
assert data == {"status": "running"}
|
||||
|
||||
|
||||
class TestScriptDeltaEvent:
|
||||
"""流式 script_delta 事件推送(worker on_delta 回调语义)测试。
|
||||
|
||||
这里不直接导入 _step_script_generation(依赖 celery/db),而是复刻 on_delta 回调
|
||||
的核心节流与事件格式逻辑,验证:
|
||||
1) 事件 type = viral_video:script_delta
|
||||
2) data 包含 delta / full_text / text_length
|
||||
3) 节流逻辑(两次事件间隔 ≥200ms)
|
||||
"""
|
||||
|
||||
def _make_on_delta(self, job_id, stage, emit_fn, start_ts):
|
||||
"""复刻 worker 里 _on_script_delta 回调的关键逻辑。"""
|
||||
import time
|
||||
|
||||
last_emit_ts = [start_ts]
|
||||
|
||||
def on_delta(delta, full_text):
|
||||
now = time.time()
|
||||
if now - last_emit_ts[0] < 0.2:
|
||||
return
|
||||
last_emit_ts[0] = now
|
||||
_est_progress = min(70.0, 15.0 + (len(full_text) / 3500.0) * 55.0)
|
||||
emit_fn(
|
||||
job_id,
|
||||
stage,
|
||||
round(_est_progress, 1),
|
||||
f"generating ({len(full_text)} chars)",
|
||||
{"delta": delta, "full_text": full_text, "text_length": len(full_text)},
|
||||
event_type="viral_video:script_delta",
|
||||
)
|
||||
|
||||
return on_delta
|
||||
|
||||
def test_script_delta_event_format(self):
|
||||
"""script_delta 事件格式与字段完整(用 time.sleep 跨过节流窗口)"""
|
||||
import time
|
||||
|
||||
events = []
|
||||
|
||||
def fake_emit(job_id, stage, progress, msg, data, event_type):
|
||||
events.append(
|
||||
{
|
||||
"job_id": job_id,
|
||||
"stage": stage,
|
||||
"progress": progress,
|
||||
"message": msg,
|
||||
"data": data,
|
||||
"type": event_type,
|
||||
}
|
||||
)
|
||||
|
||||
on_delta = self._make_on_delta("job123", "script_generation", fake_emit, time.time() - 1)
|
||||
on_delta("你好", "你好")
|
||||
time.sleep(0.25)
|
||||
on_delta("世界", "你好世界")
|
||||
|
||||
assert len(events) >= 2
|
||||
e = events[0]
|
||||
assert e["type"] == "viral_video:script_delta"
|
||||
assert e["job_id"] == "job123"
|
||||
assert e["data"]["delta"] == "你好"
|
||||
assert e["data"]["full_text"] == "你好"
|
||||
assert e["data"]["text_length"] == 2
|
||||
assert events[-1]["data"]["full_text"] == "你好世界"
|
||||
|
||||
def test_rate_limit_throttles_fast_calls(self):
|
||||
"""节流:200ms 内的连续 delta 在第一次发送后被拒绝;跨节流窗口的会放行"""
|
||||
import time
|
||||
|
||||
events = []
|
||||
|
||||
def fake_emit(*a, **kw):
|
||||
events.append(kw.get("event_type", a[5] if len(a) > 5 else "x"))
|
||||
|
||||
# start_ts 设为 1s 前,保证第一次 emit 被放行
|
||||
on_delta = self._make_on_delta("j", "s", fake_emit, time.time() - 1)
|
||||
on_delta("a", "a") # 放行:距离 start_ts 已 1s
|
||||
on_delta("b", "ab") # 被节流:距上次 <0.2s
|
||||
on_delta("c", "abc") # 被节流:同上
|
||||
time.sleep(0.25)
|
||||
on_delta("d", "abcd") # 放行:已跨过节流窗口
|
||||
assert len(events) == 2
|
||||
assert all(e == "viral_video:script_delta" for e in events)
|
||||
|
||||
Reference in New Issue
Block a user