99import logging
1010from typing import Callable , List , Tuple
1111
12+ from ..runtime_api_resolver import get_runtime_api_override
1213from .common import CommonTranslator , validate_openai_response
1314from .keys import SAKURA_API_BASE , SAKURA_DICT_PATH
1415
16+ DEFAULT_SAKURA_API_BASE = 'http://127.0.0.1:8080/v1'
17+ DEFAULT_SAKURA_DICT_PATH = './dict/sakura_dict.txt'
18+
1519
1620class SakuraDict ():
1721 def __init__ (self , path : str , logger : logging .Logger ) -> None :
@@ -209,12 +213,10 @@ class SakuraTranslator(CommonTranslator):
209213
210214 def __init__ (self ):
211215 super ().__init__ ()
212- self .client = openai .AsyncOpenAI (api_key = openai .api_key or 'empty' )
213- if "/v1" not in SAKURA_API_BASE :
214- self .client .base_url = SAKURA_API_BASE + "/v1"
215- else :
216- self .client .base_url = SAKURA_API_BASE
217- self .client .api_key = "sk-114514"
216+ self .api_base = self ._normalize_api_base (os .getenv ('SAKURA_API_BASE' ) or SAKURA_API_BASE )
217+ self .dict_path = os .getenv ('SAKURA_DICT_PATH' ) or SAKURA_DICT_PATH or DEFAULT_SAKURA_DICT_PATH
218+ self .client = None
219+ self ._setup_client ()
218220 self .temperature = 0.3
219221 self .top_p = 0.3
220222 self .frequency_penalty = 0.1
@@ -223,8 +225,56 @@ def __init__(self):
223225 self ._heart_pattern = re .compile (r'❤' )
224226 self .sakura_dict = SakuraDict (self .get_dict_path (), self .logger )
225227
228+ @staticmethod
229+ def _normalize_api_base (api_base : str ) -> str :
230+ normalized = (
231+ str (api_base or SAKURA_API_BASE or DEFAULT_SAKURA_API_BASE ).strip ()
232+ or DEFAULT_SAKURA_API_BASE
233+ )
234+ normalized = normalized .rstrip ("/" )
235+ if "/v1" not in normalized :
236+ normalized = f"{ normalized } /v1"
237+ return normalized
238+
239+ def _setup_client (self ):
240+ self .client = openai .AsyncOpenAI (
241+ api_key = "sk-114514" ,
242+ base_url = self .api_base ,
243+ )
244+
245+ def parse_args (self , config ):
246+ super ().parse_args (config )
247+ translator_args = self ._resolve_translator_config (config )
248+ user_env_vars = getattr (config , "_user_env_vars" , None ) or {}
249+ runtime_override = get_runtime_api_override (config , "translator" , "sakura" )
250+
251+ api_base = (
252+ self ._get_config_value (translator_args , "user_api_base" , None )
253+ or runtime_override .get ("api_base" )
254+ or user_env_vars .get ("SAKURA_API_BASE" )
255+ or os .getenv ("SAKURA_API_BASE" )
256+ or SAKURA_API_BASE
257+ or DEFAULT_SAKURA_API_BASE
258+ )
259+ api_base = self ._normalize_api_base (api_base )
260+ if api_base != self .api_base :
261+ self .api_base = api_base
262+ self ._setup_client ()
263+ self .logger .info (f"Sakura API base updated: { self .api_base } " )
264+
265+ dict_path = (
266+ user_env_vars .get ("SAKURA_DICT_PATH" )
267+ or os .getenv ("SAKURA_DICT_PATH" )
268+ or SAKURA_DICT_PATH
269+ or DEFAULT_SAKURA_DICT_PATH
270+ )
271+ dict_path = str (dict_path or DEFAULT_SAKURA_DICT_PATH ).strip () or DEFAULT_SAKURA_DICT_PATH
272+ if dict_path != self .dict_path :
273+ self .dict_path = dict_path
274+ self .sakura_dict = SakuraDict (self .get_dict_path (), self .logger )
275+
226276 def get_dict_path (self ):
227- return SAKURA_DICT_PATH
277+ return getattr ( self , 'dict_path' , None ) or SAKURA_DICT_PATH or DEFAULT_SAKURA_DICT_PATH
228278
229279 def detect_and_caculate_repeats (self , s : str , threshold : int = _REPEAT_DETECT_THRESHOLD , remove_all = True ) -> Tuple [bool , str , int , str ]:
230280 """
@@ -457,6 +507,7 @@ def _delete_quotation_mark(self, texts: List[str]) -> List[str]:
457507
458508 async def _translate (self , from_lang : str , to_lang : str , queries : List [str ], ctx = None ) -> List [str ]:
459509 self .logger .debug (f'Temperature: { self .temperature } , TopP: { self .top_p } ' )
510+ self .logger .info (f'Sakura当前连接地址: { self .api_base } ' )
460511 self .logger .debug (f'原文: { queries } ' )
461512 text_prompt = '\n ' .join (queries )
462513 self .logger .debug ('-- Sakura Prompt --\n ' + self ._format_prompt_log (text_prompt ) + '\n \n ' )
@@ -500,11 +551,9 @@ async def _handle_translation_request(self, prompt) -> str:
500551 except (openai .APIError , openai .APIConnectionError ) as e :
501552 server_error_attempt += 1
502553 if server_error_attempt >= self ._RETRY_ATTEMPTS :
503- self .logger .error (f'Sakura API请求失败。错误信息: { e } ' )
504- if isinstance (prompt , list ):
505- return "\n " .join (prompt )
506- return str (prompt )
507- self .logger .warning (f'Sakura因服务器错误而进行重试。尝试次数: { server_error_attempt } ,错误信息: { e } ' )
554+ self .logger .error (f'Sakura API请求失败。地址:{ self .api_base } ,错误信息: { e } ' )
555+ raise Exception (f'Sakura API请求失败(地址:{ self .api_base } ):{ e } ' ) from e
556+ self .logger .warning (f'Sakura因服务器错误而进行重试。地址:{ self .api_base } ,尝试次数: { server_error_attempt } ,错误信息: { e } ' )
508557
509558 return response
510559
0 commit comments