This commit is contained in:
jinli.yl 2026-04-08 16:49:36 +08:00
parent 16d139583f
commit 2665f31d86
8 changed files with 37 additions and 192 deletions

View file

@ -1,9 +1,7 @@
"""Module for registering AgentScope LLM models."""
from agentscope.model import DashScopeChatModel
from agentscope.model import OpenAIChatModel
from ..registry_factory import R
R.as_llms.register("openai")(OpenAIChatModel)
R.as_llms.register("dashscope")(DashScopeChatModel)

View file

@ -2,10 +2,14 @@
from abc import ABC, abstractmethod
from ..enumeration import ComponentEnum
class BaseComponent(ABC):
"""Base class supporting async start/close and async context management."""
component_type = ComponentEnum.BASE
@abstractmethod
async def start(self) -> None:
"""Start the component asynchronously."""
@ -13,13 +17,16 @@ class BaseComponent(ABC):
@abstractmethod
async def close(self) -> None:
"""Close the component asynchronously."""
...
async def __aenter__(self) -> "BaseComponent":
"""Enter async context manager."""
await self.start()
return self
async def __aexit__(self, exc_type, exc_val, exc_tb) -> None:
async def __aexit__(self, exc_type, exc_val, exc_tb) -> bool:
"""Exit async context manager."""
await self.close()
if exc_type is not None:
return True
return False

View file

@ -1,41 +0,0 @@
"""Module providing a dictionary subclass with attribute-style access and pickling support."""
from typing import Generic, TypeVar
_KT = TypeVar("_KT")
_VT = TypeVar("_VT")
class BaseDict(dict, Generic[_KT, _VT]):
"""A dictionary subclass that enables accessing and modifying keys as attributes."""
def __getattr__(self, name: str) -> _VT:
"""Retrieve a dictionary item as an attribute."""
try:
return self[name]
except KeyError as e:
raise AttributeError(f"'{type(self).__name__}' object has no attribute '{name}'") from e
def __setattr__(self, name: str, value: _VT) -> None:
"""Assign a value to a dictionary item using attribute syntax."""
self[name] = value
def __delattr__(self, name: str) -> None:
"""Remove a dictionary item using attribute syntax."""
try:
# Delete item from dict via key
del self[name]
except KeyError as e:
raise AttributeError(f"'{type(self).__name__}' object has no attribute '{name}'") from e
def __getstate__(self) -> dict:
"""Return the dictionary representation for pickling."""
return dict(self)
def __setstate__(self, state: dict) -> None:
"""Restore the dictionary state from a pickled object."""
self.update(state)
def __reduce__(self):
"""Define the reconstruction logic for pickling processes."""
return self.__class__, (), self.__getstate__()

View file

@ -1,146 +0,0 @@
"""Module for managing and formatting prompt templates from files or dictionaries."""
import json
from pathlib import Path
from string import Formatter
from typing import Any, Dict, Optional, Union
import yaml
from loguru import logger
from .base_dict import BaseDict
class PromptHandler(BaseDict):
"""A context-aware handler for loading, retrieving, and formatting prompt templates."""
def __init__(self, language: str = "", **kwargs):
super().__init__(**kwargs)
# Use object.__setattr__ to avoid storing 'language' in the dict
object.__setattr__(self, "language", language.strip())
def load_prompt_by_file(
self,
prompt_file_path: Optional[Union[Path, str]] = None,
overwrite: bool = True,
) -> "PromptHandler":
"""Load prompt configurations from a YAML or JSON file."""
if prompt_file_path is None:
return self
if isinstance(prompt_file_path, str):
prompt_file_path = Path(prompt_file_path)
if not prompt_file_path.exists():
return self
suffix = prompt_file_path.suffix.lower()
with prompt_file_path.open(encoding="utf-8") as f:
if suffix in [".yaml", ".yml"]:
prompt_dict = yaml.safe_load(f)
elif suffix == ".json":
prompt_dict = json.load(f)
else:
raise ValueError(f"Unsupported file format: {suffix}")
self.load_prompt_dict(prompt_dict, overwrite=overwrite)
return self
def load_prompt_dict(
self,
prompt_dict: Optional[Dict[str, Any]] = None,
overwrite: bool = True,
) -> "PromptHandler":
"""Merge a dictionary of prompt strings into the current context."""
if not prompt_dict:
return self
for key, value in prompt_dict.items():
if not isinstance(value, str):
continue
if key in self:
if overwrite:
logger.warning(f"Overwriting prompt '{key}'")
self[key] = value
else:
self[key] = value
return self
def get_prompt(self, prompt_name: str, fallback_to_base: bool = True) -> str:
"""Retrieve a prompt by name with automatic language suffix handling."""
if self.language and not prompt_name.endswith(f"_{self.language}"):
key_with_lang = f"{prompt_name}_{self.language}"
if key_with_lang in self:
return self[key_with_lang].strip()
if prompt_name in self:
return self[prompt_name].strip()
if fallback_to_base and self.language and prompt_name.endswith(f"_{self.language}"):
base_name = prompt_name[: -(len(self.language) + 1)]
if base_name in self:
return self[base_name].strip()
raise KeyError(f"Prompt '{prompt_name}' not found. Available: {list(self.keys())[:10]}")
def has_prompt(self, prompt_name: str) -> bool:
"""Check if a prompt exists."""
try:
self.get_prompt(prompt_name)
return True
except KeyError:
return False
def list_prompts(self, language_filter: Optional[str] = None) -> list[str]:
"""List all available prompt names."""
if language_filter is None:
return list(self.keys())
suffix = f"_{language_filter.strip()}"
return [key for key in self.keys() if key.endswith(suffix)]
@staticmethod
def _extract_format_fields(template: str) -> set[str]:
"""Extract all format field names from a template string."""
return {field_name for _, field_name, _, _ in Formatter().parse(template) if field_name is not None}
@staticmethod
def _filter_conditional_lines(prompt: str, flags: Dict[str, bool]) -> str:
"""Filter lines based on boolean flags."""
filtered_lines = []
for line in prompt.split("\n"):
matched_flag = None
for flag_name in flags:
if line.startswith(f"[{flag_name}]"):
matched_flag = flag_name
break
if matched_flag is None:
filtered_lines.append(line)
elif flags[matched_flag]:
filtered_lines.append(line[len(f"[{matched_flag}]") :])
return "\n".join(filtered_lines)
def prompt_format(self, prompt_name: str, validate: bool = True, **kwargs) -> str:
"""Format a prompt with conditional line filtering and variable substitution."""
prompt = self.get_prompt(prompt_name)
flag_kwargs = {k: v for k, v in kwargs.items() if isinstance(v, bool)}
format_kwargs = {k: v for k, v in kwargs.items() if not isinstance(v, bool)}
if flag_kwargs:
prompt = self._filter_conditional_lines(prompt, flag_kwargs)
if validate:
required_fields = self._extract_format_fields(prompt)
missing_fields = required_fields - set(format_kwargs.keys())
if missing_fields:
raise ValueError(f"Missing format variables for '{prompt_name}': {sorted(missing_fields)}")
if format_kwargs:
prompt = prompt.format(**format_kwargs)
return prompt.strip()
def __repr__(self) -> str:
return f"PromptHandler(language='{self.language}', num_prompts={len(self)})"

View file

@ -4,7 +4,7 @@ import inspect
from typing import Callable, TypeVar
from .base_dict import BaseDict
from .utils import singleton
from ..utils import singleton
T = TypeVar("T")

View file

@ -0,0 +1,7 @@
"""enumeration"""
from .component_enum import ComponentEnum
__all__ = [
"ComponentEnum",
]

View file

@ -0,0 +1,15 @@
from enum import Enum
class ComponentEnum(str, Enum):
BASE = "base"
AS_LLM = "as_llm"
AS_LLM_FORMATTER = "as_llm_formatter"
FILE_STORE = "file_store"
FILE_WATCHER = "file_watcher"

View file

@ -0,0 +1,5 @@
from .singleton import singleton
__all__ = [
"singleton",
]