mirror of
https://github.com/agentscope-ai/ReMe.git
synced 2026-10-08 03:10:24 +00:00
init
This commit is contained in:
parent
16d139583f
commit
2665f31d86
8 changed files with 37 additions and 192 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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__()
|
||||
|
|
@ -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)})"
|
||||
|
|
@ -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")
|
||||
|
||||
|
|
|
|||
7
reme_cli/enumeration/__init__.py
Normal file
7
reme_cli/enumeration/__init__.py
Normal file
|
|
@ -0,0 +1,7 @@
|
|||
"""enumeration"""
|
||||
|
||||
from .component_enum import ComponentEnum
|
||||
|
||||
__all__ = [
|
||||
"ComponentEnum",
|
||||
]
|
||||
15
reme_cli/enumeration/component_enum.py
Normal file
15
reme_cli/enumeration/component_enum.py
Normal 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"
|
||||
|
|
@ -0,0 +1,5 @@
|
|||
from .singleton import singleton
|
||||
|
||||
__all__ = [
|
||||
"singleton",
|
||||
]
|
||||
Loading…
Add table
Reference in a new issue