v1.0
This commit is contained in:
@@ -0,0 +1,108 @@
|
||||
import json
|
||||
import time
|
||||
from typing import Dict, Any, Optional
|
||||
|
||||
import requests
|
||||
|
||||
|
||||
class APIClient:
|
||||
"""DeepSeek API 客户端,含重试逻辑。"""
|
||||
|
||||
def __init__(self, api_key: str, model: str = 'deepseek-v4-flash',
|
||||
base_url: str = 'https://api.deepseek.com/chat/completions',
|
||||
max_retries: int = 3, timeout: int = 120):
|
||||
self.api_key = api_key
|
||||
self.model = model
|
||||
self.base_url = base_url
|
||||
self.max_retries = max_retries
|
||||
self.timeout = timeout
|
||||
|
||||
def generate(self, prompt: str) -> Dict[str, Any]:
|
||||
"""发送 prompt 并返回 AI 生成的结果。"""
|
||||
system_prompt = (
|
||||
"你是COBOL程序的测试数据生成专家。"
|
||||
"请严格按照提供的规则,生成符合格式要求的测试数据。"
|
||||
"输出必须是可被json.loads()直接解析的JSON,不要包裹在```json```代码块中。"
|
||||
"不要在JSON前后添加任何说明文字。"
|
||||
)
|
||||
|
||||
messages = [
|
||||
{"role": "system", "content": system_prompt},
|
||||
{"role": "user", "content": prompt},
|
||||
]
|
||||
|
||||
last_error = None
|
||||
|
||||
for attempt in range(1, self.max_retries + 1):
|
||||
try:
|
||||
response = self._call_api(messages)
|
||||
content = response['choices'][0]['message']['content']
|
||||
data = self._parse_json(content)
|
||||
if data is not None:
|
||||
return data
|
||||
|
||||
error_msg = (
|
||||
f"前回の出力は有効なJSONではありませんでした。"
|
||||
f"必ず有効なJSONのみを出力してください。"
|
||||
f"コードブロック(```)で囲まないでください。"
|
||||
)
|
||||
messages.append({"role": "assistant", "content": content})
|
||||
messages.append({"role": "user", "content": error_msg})
|
||||
|
||||
except requests.exceptions.RequestException as e:
|
||||
last_error = e
|
||||
if attempt < self.max_retries:
|
||||
time.sleep(2 ** attempt)
|
||||
continue
|
||||
|
||||
raise RuntimeError(
|
||||
f"API调用失败,已重试{self.max_retries}次。"
|
||||
f"最后错误: {last_error}"
|
||||
)
|
||||
|
||||
def _call_api(self, messages: list) -> dict:
|
||||
"""单次 API 调用。"""
|
||||
headers = {
|
||||
'Authorization': f'Bearer {self.api_key}',
|
||||
'Content-Type': 'application/json',
|
||||
}
|
||||
payload = {
|
||||
'model': self.model,
|
||||
'messages': messages,
|
||||
'temperature': 0.3,
|
||||
'max_tokens': 8192,
|
||||
}
|
||||
|
||||
resp = requests.post(
|
||||
self.base_url,
|
||||
headers=headers,
|
||||
json=payload,
|
||||
timeout=self.timeout
|
||||
)
|
||||
resp.raise_for_status()
|
||||
return resp.json()
|
||||
|
||||
@staticmethod
|
||||
def _parse_json(text: str) -> Optional[dict]:
|
||||
"""尝试从文本中提取 JSON。"""
|
||||
text = text.strip()
|
||||
|
||||
if text.startswith('```json'):
|
||||
text = text[7:]
|
||||
if text.startswith('```'):
|
||||
text = text[3:]
|
||||
if text.endswith('```'):
|
||||
text = text[:-3]
|
||||
text = text.strip()
|
||||
|
||||
try:
|
||||
return json.loads(text)
|
||||
except json.JSONDecodeError:
|
||||
start = text.find('{')
|
||||
end = text.rfind('}')
|
||||
if start >= 0 and end > start:
|
||||
try:
|
||||
return json.loads(text[start:end + 1])
|
||||
except json.JSONDecodeError:
|
||||
pass
|
||||
return None
|
||||
Reference in New Issue
Block a user