Spaces:
Runtime error
Runtime error
Download backend/verify_backend_complete.py from cytopa99/universal-fast-dubbing: direct link, hf CLI and curl.
- Browser
- Download file 23.6 kB
-
https://huggingface.co/spaces/cytopa99/universal-fast-dubbing/resolve/main/backend/verify_backend_complete.py
- Command line
-
hf download hf://spaces/cytopa99/universal-fast-dubbing/backend/verify_backend_complete.py
-
curl -L -o verify_backend_complete.py https://huggingface.co/spaces/cytopa99/universal-fast-dubbing/resolve/main/backend/verify_backend_complete.py
23.6 kB
| """ | |
| 后端完整性验证脚本 | |
| 验证所有后端模块是否正常集成工作,包括: | |
| - 核心处理模块 | |
| - 错误处理模块 | |
| - 日志记录模块 | |
| - 并行处理池 | |
| - 请求路由器 | |
| - API网关 | |
| Requirements: 检查点任务 11 | |
| """ | |
| import sys | |
| import os | |
| import asyncio | |
| from typing import Dict, List, Tuple | |
| sys.path.insert(0, '.') | |
| def print_header(title: str) -> None: | |
| """打印标题""" | |
| print("\n" + "=" * 60) | |
| print(f" {title}") | |
| print("=" * 60) | |
| def print_section(title: str) -> None: | |
| """打印章节标题""" | |
| print(f"\n--- {title} ---") | |
| def check_result(name: str, passed: bool, details: str = "") -> Tuple[str, bool]: | |
| """打印检查结果""" | |
| status = "✓" if passed else "✗" | |
| detail_str = f" ({details})" if details else "" | |
| print(f" {status} {name}{detail_str}") | |
| return (name, passed) | |
| class BackendVerifier: | |
| """后端完整性验证器""" | |
| def __init__(self): | |
| self.results: List[Tuple[str, bool]] = [] | |
| def verify_imports(self) -> bool: | |
| """验证所有模块可以正确导入""" | |
| print_section("模块导入检查") | |
| all_passed = True | |
| # 错误处理模块 | |
| try: | |
| from modules import ( | |
| ErrorCode, ErrorType, ErrorResponse, ErrorFactory, | |
| create_error_response | |
| ) | |
| self.results.append(check_result("错误处理模块", True)) | |
| except Exception as e: | |
| self.results.append(check_result("错误处理模块", False, str(e))) | |
| all_passed = False | |
| # 日志模块 | |
| try: | |
| from modules import ( | |
| LogLevel, Component, StructuredLogRecord, ComponentLogger, | |
| setup_logging, get_component_logger, log_performance | |
| ) | |
| self.results.append(check_result("日志记录模块", True)) | |
| except Exception as e: | |
| self.results.append(check_result("日志记录模块", False, str(e))) | |
| all_passed = False | |
| # Groq客户端 | |
| try: | |
| from modules import ( | |
| GroqClient, GroqConfig, GroqError, GroqRateLimitError, | |
| GroqTimeoutError, GroqAuthError, GroqConnectionError, RetryStats | |
| ) | |
| self.results.append(check_result("Groq客户端模块", True)) | |
| except Exception as e: | |
| self.results.append(check_result("Groq客户端模块", False, str(e))) | |
| all_passed = False | |
| # TTS生成器 | |
| try: | |
| from modules import ( | |
| TTSGenerator, TTSConfig, TTSError, VoiceRole | |
| ) | |
| self.results.append(check_result("TTS生成器模块", True)) | |
| except Exception as e: | |
| self.results.append(check_result("TTS生成器模块", False, str(e))) | |
| all_passed = False | |
| # 智能分段器 | |
| try: | |
| from modules import ( | |
| SmartSegmenter, SegmenterConfig, SegmenterError, SegmentInfo | |
| ) | |
| self.results.append(check_result("智能分段器模块", True)) | |
| except Exception as e: | |
| self.results.append(check_result("智能分段器模块", False, str(e))) | |
| all_passed = False | |
| # 音频同步引擎 | |
| try: | |
| from modules import ( | |
| AudioSyncEngine, SyncConfig, AudioSyncError | |
| ) | |
| self.results.append(check_result("音频同步引擎模块", True)) | |
| except Exception as e: | |
| self.results.append(check_result("音频同步引擎模块", False, str(e))) | |
| all_passed = False | |
| # 并行处理池 | |
| try: | |
| from modules import ( | |
| ParallelProcessingPool, ProcessorConfig, ProcessingError, | |
| SegmentResult | |
| ) | |
| self.results.append(check_result("并行处理池模块", True)) | |
| except Exception as e: | |
| self.results.append(check_result("并行处理池模块", False, str(e))) | |
| all_passed = False | |
| # 请求路由器 | |
| try: | |
| from modules import ( | |
| RequestRouter, RouterConfig, RouterError, URLNotSupportedError, | |
| DownloadError, ProcessingMode | |
| ) | |
| self.results.append(check_result("请求路由器模块", True)) | |
| except Exception as e: | |
| self.results.append(check_result("请求路由器模块", False, str(e))) | |
| all_passed = False | |
| # API网关 | |
| try: | |
| from modules import ( | |
| GradioAPIGateway, GatewayConfig, GatewayError, CacheEntry | |
| ) | |
| self.results.append(check_result("API网关模块", True)) | |
| except Exception as e: | |
| self.results.append(check_result("API网关模块", False, str(e))) | |
| all_passed = False | |
| return all_passed | |
| def verify_error_handling(self) -> bool: | |
| """验证错误处理功能""" | |
| print_section("错误处理功能检查") | |
| all_passed = True | |
| try: | |
| from modules import ( | |
| ErrorCode, ErrorType, ErrorResponse, ErrorFactory | |
| ) | |
| # 测试错误响应创建 | |
| error = ErrorResponse( | |
| error_code=ErrorCode.GROQ_RATE_LIMIT, | |
| error_type=ErrorType.RETRYABLE, | |
| message="测试限流错误", | |
| retry_available=True, | |
| suggested_action="等待30秒后重试" | |
| ) | |
| # 验证to_dict | |
| error_dict = error.to_dict() | |
| assert "error_code" in error_dict | |
| assert "message" in error_dict | |
| self.results.append(check_result("ErrorResponse.to_dict()", True)) | |
| # 验证to_user_message | |
| user_msg = error.to_user_message() | |
| assert "测试限流错误" in user_msg | |
| self.results.append(check_result("ErrorResponse.to_user_message()", True)) | |
| # 测试工厂方法 | |
| rate_limit = ErrorFactory.create_groq_rate_limit_error(retry_after=30) | |
| assert rate_limit.error_code == ErrorCode.GROQ_RATE_LIMIT | |
| self.results.append(check_result("ErrorFactory.create_groq_rate_limit_error()", True)) | |
| timeout = ErrorFactory.create_groq_timeout_error(timeout=30, operation="语音识别") | |
| assert timeout.error_code == ErrorCode.GROQ_TIMEOUT | |
| self.results.append(check_result("ErrorFactory.create_groq_timeout_error()", True)) | |
| url_error = ErrorFactory.create_url_not_supported_error(url="https://test.com") | |
| assert url_error.error_code == ErrorCode.URL_NOT_SUPPORTED | |
| self.results.append(check_result("ErrorFactory.create_url_not_supported_error()", True)) | |
| except Exception as e: | |
| self.results.append(check_result("错误处理功能", False, str(e))) | |
| all_passed = False | |
| return all_passed | |
| def verify_logging(self) -> bool: | |
| """验证日志记录功能""" | |
| print_section("日志记录功能检查") | |
| all_passed = True | |
| try: | |
| from modules import ( | |
| Component, StructuredLogRecord, ComponentLogger, | |
| get_component_logger, StructuredFormatter, HumanReadableFormatter | |
| ) | |
| # 测试结构化日志记录 | |
| record = StructuredLogRecord( | |
| timestamp="2024-01-01T00:00:00", | |
| level="INFO", | |
| component="Test", | |
| message="测试消息", | |
| session_id="test-session", | |
| duration_ms=100.5 | |
| ) | |
| # 验证to_dict | |
| record_dict = record.to_dict() | |
| assert record_dict["component"] == "Test" | |
| assert record_dict["duration_ms"] == 100.5 | |
| self.results.append(check_result("StructuredLogRecord.to_dict()", True)) | |
| # 验证to_json | |
| json_str = record.to_json() | |
| assert "Test" in json_str | |
| assert "测试消息" in json_str | |
| self.results.append(check_result("StructuredLogRecord.to_json()", True)) | |
| # 测试组件日志记录器 | |
| logger = get_component_logger(Component.GROQ_CLIENT) | |
| assert logger.component == "GroqClient" | |
| self.results.append(check_result("get_component_logger()", True)) | |
| # 测试格式化器 | |
| formatter = StructuredFormatter() | |
| assert formatter is not None | |
| self.results.append(check_result("StructuredFormatter", True)) | |
| human_formatter = HumanReadableFormatter(use_colors=False) | |
| assert human_formatter is not None | |
| self.results.append(check_result("HumanReadableFormatter", True)) | |
| except Exception as e: | |
| self.results.append(check_result("日志记录功能", False, str(e))) | |
| all_passed = False | |
| return all_passed | |
| def verify_groq_client(self) -> bool: | |
| """验证Groq客户端功能""" | |
| print_section("Groq客户端功能检查") | |
| all_passed = True | |
| try: | |
| from modules import GroqClient, GroqConfig, RetryStats | |
| # 创建配置 | |
| config = GroqConfig( | |
| api_key="test_key", | |
| max_retries=3, | |
| retry_base_delay=1.0, | |
| retry_max_delay=30.0, | |
| retry_jitter=True | |
| ) | |
| # 创建客户端 | |
| client = GroqClient(config) | |
| # 验证配置 | |
| assert client.asr_model == "whisper-large-v3-turbo" | |
| assert client.llm_model == "llama3-8b-8192" | |
| self.results.append(check_result("GroqClient配置", True)) | |
| # 验证重试统计 | |
| stats = client.get_retry_stats() | |
| assert isinstance(stats, list) | |
| self.results.append(check_result("GroqClient.get_retry_stats()", True)) | |
| # 验证退避延迟计算 | |
| delay0 = client._calculate_backoff_delay(0) | |
| delay1 = client._calculate_backoff_delay(1) | |
| delay2 = client._calculate_backoff_delay(2) | |
| # 指数退避:delay应该递增 | |
| assert delay1 > delay0 * 0.8 # 考虑抖动 | |
| assert delay2 > delay1 * 0.8 | |
| self.results.append(check_result("指数退避延迟计算", True)) | |
| # 验证retry_after优先 | |
| delay_with_retry = client._calculate_backoff_delay(0, retry_after=10.0) | |
| assert delay_with_retry == 10.0 | |
| self.results.append(check_result("retry_after优先级", True)) | |
| except Exception as e: | |
| self.results.append(check_result("Groq客户端功能", False, str(e))) | |
| all_passed = False | |
| return all_passed | |
| def verify_tts_generator(self) -> bool: | |
| """验证TTS生成器功能""" | |
| print_section("TTS生成器功能检查") | |
| all_passed = True | |
| try: | |
| from modules import TTSGenerator, VoiceRole | |
| # 创建生成器 | |
| generator = TTSGenerator() | |
| # 验证语音映射 | |
| voices = generator.get_available_voices() | |
| assert "MALE" in voices | |
| assert "FEMALE" in voices | |
| assert "CHILD" in voices | |
| assert "NARRATOR" in voices | |
| self.results.append(check_result("语音角色映射", True)) | |
| # 验证语音模型 | |
| assert voices["MALE"] == "zh-CN-YunxiNeural" | |
| assert voices["FEMALE"] == "zh-CN-XiaoxiaoNeural" | |
| assert voices["CHILD"] == "zh-CN-YunjianNeural" | |
| assert voices["NARRATOR"] == "zh-CN-YunyangNeural" | |
| self.results.append(check_result("语音模型配置", True)) | |
| # 验证VoiceRole枚举 | |
| assert VoiceRole.MALE.value == "MALE" | |
| assert VoiceRole.FEMALE.value == "FEMALE" | |
| self.results.append(check_result("VoiceRole枚举", True)) | |
| except Exception as e: | |
| self.results.append(check_result("TTS生成器功能", False, str(e))) | |
| all_passed = False | |
| return all_passed | |
| def verify_segmenter(self) -> bool: | |
| """验证智能分段器功能""" | |
| print_section("智能分段器功能检查") | |
| all_passed = True | |
| try: | |
| from modules import SmartSegmenter, SegmenterConfig, SegmentInfo | |
| # 创建分段器 | |
| segmenter = SmartSegmenter() | |
| # 验证配置 | |
| assert segmenter.config.max_segment_duration == 480.0 # 8分钟 | |
| assert segmenter.config.min_segment_duration == 300.0 # 5分钟 | |
| assert segmenter.config.silence_threshold_db == -40.0 | |
| self.results.append(check_result("分段器配置", True)) | |
| # 验证should_segment方法存在 | |
| assert hasattr(segmenter, 'should_segment') | |
| assert hasattr(segmenter, 'segment_audio') | |
| self.results.append(check_result("分段器方法", True)) | |
| # 验证SegmentInfo数据类 | |
| segment = SegmentInfo( | |
| index=0, | |
| start_time=0.0, | |
| end_time=300.0, | |
| duration=300.0, | |
| audio_path="test.wav" | |
| ) | |
| assert segment.duration == 300.0 | |
| self.results.append(check_result("SegmentInfo数据类", True)) | |
| except Exception as e: | |
| self.results.append(check_result("智能分段器功能", False, str(e))) | |
| all_passed = False | |
| return all_passed | |
| def verify_audio_sync(self) -> bool: | |
| """验证音频同步引擎功能""" | |
| print_section("音频同步引擎功能检查") | |
| all_passed = True | |
| try: | |
| from modules import AudioSyncEngine, SyncConfig | |
| # 创建同步引擎 | |
| engine = AudioSyncEngine() | |
| # 验证配置 | |
| assert engine.config.max_speed_ratio == 1.4 | |
| assert engine.config.sync_tolerance == 0.3 | |
| self.results.append(check_result("同步引擎配置", True)) | |
| # 验证方法存在 | |
| assert hasattr(engine, 'align') | |
| assert hasattr(engine, 'align_segment') | |
| assert hasattr(engine, 'check_sync_drift') | |
| self.results.append(check_result("同步引擎方法", True)) | |
| # 测试同步漂移检查(返回元组: (needs_correction, drift)) | |
| needs_correction, drift = engine.check_sync_drift(10.0, 10.2) | |
| assert isinstance(drift, float) | |
| assert drift < 0.3 # 偏差小于容差 | |
| assert needs_correction == False # 不需要校正 | |
| self.results.append(check_result("同步漂移检查", True)) | |
| except Exception as e: | |
| self.results.append(check_result("音频同步引擎功能", False, str(e))) | |
| all_passed = False | |
| return all_passed | |
| def verify_processor(self) -> bool: | |
| """验证并行处理池功能""" | |
| print_section("并行处理池功能检查") | |
| all_passed = True | |
| try: | |
| from modules import ParallelProcessingPool, ProcessorConfig, SegmentResult | |
| # 创建处理池(不初始化,只检查结构) | |
| config = ProcessorConfig(max_workers=3) | |
| assert config.max_workers == 3 | |
| self.results.append(check_result("处理池配置", True)) | |
| # 验证SegmentResult数据类 | |
| result = SegmentResult( | |
| index=0, | |
| success=True, | |
| audio_path="test.wav", | |
| duration=100.0 | |
| ) | |
| assert result.success == True | |
| self.results.append(check_result("SegmentResult数据类", True)) | |
| except Exception as e: | |
| self.results.append(check_result("并行处理池功能", False, str(e))) | |
| all_passed = False | |
| return all_passed | |
| def verify_router(self) -> bool: | |
| """验证请求路由器功能""" | |
| print_section("请求路由器功能检查") | |
| all_passed = True | |
| try: | |
| from modules import RequestRouter, ProcessingMode | |
| # 创建路由器 | |
| router = RequestRouter() | |
| # 验证处理模式枚举 | |
| assert ProcessingMode.URL.value == "url" | |
| assert ProcessingMode.RECORD.value == "record" | |
| self.results.append(check_result("ProcessingMode枚举", True)) | |
| # 验证URL检测方法 | |
| assert hasattr(router, 'detect_platform') | |
| assert hasattr(router, 'route_request') | |
| self.results.append(check_result("路由器方法", True)) | |
| # 测试平台检测(返回元组: (platform, supports_url)) | |
| platform, supports_url = router.detect_platform("https://www.youtube.com/watch?v=test") | |
| assert platform == "youtube" | |
| assert supports_url == True | |
| self.results.append(check_result("YouTube平台检测", True)) | |
| platform, supports_url = router.detect_platform("https://www.bilibili.com/video/BV123") | |
| assert platform == "bilibili" | |
| assert supports_url == True | |
| self.results.append(check_result("Bilibili平台检测", True)) | |
| # 测试仅录制模式平台 | |
| platform, supports_url = router.detect_platform("https://www.netflix.com/watch/123") | |
| assert platform == "netflix" | |
| assert supports_url == False | |
| self.results.append(check_result("Netflix录制模式检测", True)) | |
| except Exception as e: | |
| self.results.append(check_result("请求路由器功能", False, str(e))) | |
| all_passed = False | |
| return all_passed | |
| def verify_gateway(self) -> bool: | |
| """验证API网关功能""" | |
| print_section("API网关功能检查") | |
| all_passed = True | |
| try: | |
| from modules import GradioAPIGateway, GatewayConfig, CacheEntry | |
| from datetime import datetime, timedelta | |
| # 验证配置 | |
| config = GatewayConfig( | |
| cache_duration=3600, | |
| max_sessions=10 | |
| ) | |
| assert config.cache_duration == 3600 | |
| self.results.append(check_result("网关配置", True)) | |
| # 验证CacheEntry | |
| now = datetime.now() | |
| cache = CacheEntry( | |
| result={"test": "data"}, | |
| created_at=now, | |
| expires_at=now + timedelta(hours=1) | |
| ) | |
| assert not cache.is_expired() | |
| self.results.append(check_result("CacheEntry功能", True)) | |
| # 验证过期检测 | |
| expired_cache = CacheEntry( | |
| result={"test": "data"}, | |
| created_at=now - timedelta(hours=2), | |
| expires_at=now - timedelta(hours=1) | |
| ) | |
| assert expired_cache.is_expired() | |
| self.results.append(check_result("缓存过期检测", True)) | |
| except Exception as e: | |
| self.results.append(check_result("API网关功能", False, str(e))) | |
| all_passed = False | |
| return all_passed | |
| def verify_error_integration(self) -> bool: | |
| """验证错误处理集成""" | |
| print_section("错误处理集成检查") | |
| all_passed = True | |
| try: | |
| from modules import ( | |
| ErrorFactory, GroqRateLimitError, GroqTimeoutError, | |
| URLNotSupportedError | |
| ) | |
| # 测试从异常创建错误响应 | |
| rate_limit_exc = GroqRateLimitError(retry_after=30) | |
| error_response = ErrorFactory.from_exception(rate_limit_exc) | |
| assert error_response.retry_available == True | |
| self.results.append(check_result("GroqRateLimitError转换", True)) | |
| timeout_exc = GroqTimeoutError(timeout=30, operation="测试") | |
| error_response = ErrorFactory.from_exception(timeout_exc) | |
| assert "超时" in error_response.message | |
| self.results.append(check_result("GroqTimeoutError转换", True)) | |
| except Exception as e: | |
| self.results.append(check_result("错误处理集成", False, str(e))) | |
| all_passed = False | |
| return all_passed | |
| def run_all_verifications(self) -> bool: | |
| """运行所有验证""" | |
| print_header("Universal Fast Dubbing - 后端完整性验证") | |
| # 运行各项验证 | |
| self.verify_imports() | |
| self.verify_error_handling() | |
| self.verify_logging() | |
| self.verify_groq_client() | |
| self.verify_tts_generator() | |
| self.verify_segmenter() | |
| self.verify_audio_sync() | |
| self.verify_processor() | |
| self.verify_router() | |
| self.verify_gateway() | |
| self.verify_error_integration() | |
| # 汇总结果 | |
| print_header("验证结果汇总") | |
| passed = sum(1 for _, p in self.results if p) | |
| failed = sum(1 for _, p in self.results if not p) | |
| total = len(self.results) | |
| print(f"\n 总计: {total} 项检查") | |
| print(f" 通过: {passed} 项") | |
| print(f" 失败: {failed} 项") | |
| print(f" 通过率: {passed/total*100:.1f}%") | |
| if failed > 0: | |
| print("\n 失败项目:") | |
| for name, p in self.results: | |
| if not p: | |
| print(f" ✗ {name}") | |
| print("\n" + "=" * 60) | |
| if failed == 0: | |
| print(" ✓ 后端完整性验证通过!") | |
| else: | |
| print(" ✗ 后端完整性验证失败,请检查上述问题") | |
| print("=" * 60) | |
| return failed == 0 | |
| def main(): | |
| """主函数""" | |
| verifier = BackendVerifier() | |
| success = verifier.run_all_verifications() | |
| return 0 if success else 1 | |
| if __name__ == "__main__": | |
| sys.exit(main()) | |