107 lines
2.6 KiB
Python
107 lines
2.6 KiB
Python
import threading
|
|
from typing import Any
|
|
|
|
from sympy.physics.units import temperature
|
|
|
|
from . import config
|
|
|
|
|
|
_shared_model = None
|
|
_lock = threading.Lock()
|
|
|
|
|
|
|
|
|
|
def _build_openai():
|
|
from strands.models.openai import OpenAIModel
|
|
|
|
client_args:dict[str, Any] = {"api_key": config.OPENAI_API_KEY}
|
|
if config.OPENAI_API_BASE_URL:
|
|
client_args["base_url"] = config.OPENAI_API_BASE_URL
|
|
|
|
return OpenAIModel(
|
|
client_args=client_args,
|
|
model_id=config.MODEL_ID,
|
|
params={"max_tokens": config.MAX_TOKENS, "temperature": config.TEMPERATURE},
|
|
)
|
|
|
|
|
|
def _build_anthropic():
|
|
from strands.models.anthropic import AnthropicModel
|
|
|
|
return AnthropicModel(
|
|
client_args={"api_key": config.ANTHROPIC_API_KEY},
|
|
model_id=config.MODEL_ID,
|
|
max_tokens=config.MAX_TOKENS,
|
|
params={"temperature": config.TEMPERATURE},
|
|
)
|
|
|
|
|
|
def _build_litellm():
|
|
from strands.models.litellm import LiteLLMModel
|
|
|
|
return LiteLLMModel(
|
|
model_id=config.MODEL_ID,
|
|
|
|
params={"max_tokens": config.MAX_TOKENS, "temperature": config.TEMPERATURE},
|
|
)
|
|
|
|
|
|
def _build_ollama():
|
|
from strands.models.ollama import OllamaModel
|
|
|
|
return OllamaModel(
|
|
host=config.OLLAMA_HOST,
|
|
model_id=config.MODEL_ID,
|
|
params={"temperature": config.TEMPERATURE},
|
|
)
|
|
|
|
|
|
def _build_bedrock():
|
|
from strands.models.bedrock import BedrockModel
|
|
|
|
return BedrockModel(
|
|
model_id=config.MODEL_ID,
|
|
region=config.AWS_REGION,
|
|
max_tokens=config.MAX_TOKENS,
|
|
temperature=config.TEMPERATURE,
|
|
)
|
|
|
|
|
|
_BUILDERS = {
|
|
"openai": _build_openai,
|
|
"anthropic": _build_anthropic,
|
|
"litellm": _build_litellm,
|
|
"ollama": _build_ollama,
|
|
"bedrock": _build_bedrock,
|
|
}
|
|
|
|
|
|
def get_model():
|
|
"""获取共享的模型实例,根据配置选择不同的模型提供者。
|
|
|
|
Returns:
|
|
Any: 模型实例。
|
|
"""
|
|
global _shared_model
|
|
if _shared_model is not None:
|
|
return _shared_model
|
|
|
|
with _lock:
|
|
if _shared_model is not None:
|
|
return _shared_model
|
|
|
|
builder = _BUILDERS.get(config.MODEL_PROVIDER)
|
|
if builder is None:
|
|
raise ValueError(f"Unsupported model provider: {config.MODEL_PROVIDER}"
|
|
f"Choose one of: {', '.join(_BUILDERS.keys())}")
|
|
_shared_model = builder()
|
|
return _shared_model
|
|
|
|
|
|
def reset_model() -> None:
|
|
"""重置共享的模型实例,通常用于测试或重新加载配置。"""
|
|
global _shared_model
|
|
with _lock:
|
|
_shared_model = None
|