Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2,422 changes: 872 additions & 1,550 deletions pdm.lock

Large diffs are not rendered by default.

13 changes: 6 additions & 7 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -7,12 +7,12 @@ authors = [
{"name" = "KomoriDev", "email" = "mute231010@gmail.com"}
]
dependencies = [
"arclet.entari>=0.18.0",
"arclet.entari[yaml,cron,reload,dotenv]>=0.18.3",
"agno[sqlite]>=2.7.2",
"docstring-parser>=0.17.0",
"litellm>=1.83.7",
"entari-plugin-database>=0.3.2",
"litellm>=1.84.0",
]
requires-python = ">=3.10"
requires-python = ">=3.10,<3.14"
readme = "README.md"
license = {"text" = "MIT"}

Expand All @@ -21,17 +21,16 @@ browser = [
"entari-plugin-browser>=0.5.4",
]
google = [
"litellm[google]>=1.83.7",
"litellm[google]>=1.84.0",
]

[dependency-groups]
dev = [
"arclet.entari[yaml,cron,reload,dotenv]>=0.18.0",
"arclet.entari[yaml,cron,reload,dotenv]>=0.18.3",
"entari-plugin-server>=0.7.1",
"satori-python-adapter-onebot11>=0.4.2",
"ruff>=0.14.10",
"satori-python-adapter-console>=0.5.1",
"agno>=2.5.8",
"ddgs>=9.11.2",
"satori-python-adapter-milky>=0.3.1",
]
Expand Down
2 changes: 2 additions & 0 deletions src/entari_plugin_llm/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -30,10 +30,12 @@
from .handlers import chat as chat
from .handlers import check as check
from .handlers import command as command
from .response import GenericResponse as GenericResponse
from .service import llm as llm

__all__ = [
"llm",
"LLMToolEvent",
"LLMCollectVariableEvent",
"GenericResponse",
]
114 changes: 90 additions & 24 deletions src/entari_plugin_llm/_jsondata.py
Original file line number Diff line number Diff line change
@@ -1,20 +1,53 @@
import json
from dataclasses import asdict, dataclass
from dataclasses import asdict, dataclass, field
from pathlib import Path
from typing import Any

from arclet.entari import local_data


@dataclass(slots=True)
class SessionPointer:
session_id: str
created_at: int


@dataclass(slots=True)
class LLMState:
default_model: str | None = None
current_sessions: dict[str, SessionPointer] = field(default_factory=dict)
session_models: dict[str, str] = field(default_factory=dict)

@classmethod
def from_dict(cls, data: dict[str, Any]) -> "LLMState":
value = data.get("default_model")
default_model = value if isinstance(value, str) and value else None
return cls(default_model=default_model)

current_sessions: dict[str, SessionPointer] = {}
raw_sessions = data.get("current_sessions")
if isinstance(raw_sessions, dict):
for user_id, raw_pointer in raw_sessions.items():
if not isinstance(user_id, str) or not isinstance(raw_pointer, dict):
continue
session_id = raw_pointer.get("session_id")
created_at = raw_pointer.get("created_at")
if isinstance(session_id, str) and isinstance(created_at, int):
current_sessions[user_id] = SessionPointer(session_id, created_at)

session_models: dict[str, str] = {}
raw_models = data.get("session_models")
if isinstance(raw_models, dict):
session_models = {
session_id: model
for session_id, model in raw_models.items()
if isinstance(session_id, str) and isinstance(model, str) and model
}

return cls(
default_model=default_model,
current_sessions=current_sessions,
session_models=session_models,
)

def to_dict(self) -> dict[str, Any]:
return asdict(self)
Expand All @@ -24,36 +57,37 @@ def _state_path() -> Path:
return local_data.get_data_file("entari_plugin_llm", "state.json")


def _read_state(channel: str) -> LLMState:
def _load_data() -> dict[str, Any]:
path = _state_path()
if not path.exists():
return LLMState()
return {}
try:
data = json.loads(path.read_text(encoding="utf-8"))
except (OSError, json.JSONDecodeError):
return LLMState()
if not isinstance(data, dict):
return LLMState()
res = data.get(channel)
if not isinstance(res, dict):
res = data.get("$default", {})
return LLMState.from_dict(res)
return {}
return data if isinstance(data, dict) else {}


def _read_state(channel: str) -> LLMState:
data = _load_data()
result = data.get(channel)
if not isinstance(result, dict):
result = data.get("$default")
if not isinstance(result, dict) and any(
key in data for key in ("default_model", "current_sessions", "session_models")
):
result = data
return LLMState.from_dict(result if isinstance(result, dict) else {})


def _write_state(data: LLMState, channel: str) -> None:
def _write_state(state: LLMState, channel: str) -> None:
path = _state_path()
if not path.exists():
path.parent.mkdir(parents=True, exist_ok=True)
path.write_text(json.dumps({channel: data.to_dict()}, ensure_ascii=False, indent=2), encoding="utf-8")
else:
try:
existing_data = json.loads(path.read_text(encoding="utf-8"))
except (OSError, json.JSONDecodeError):
existing_data = {}
if not isinstance(existing_data, dict):
existing_data = {}
existing_data[channel] = data.to_dict()
path.write_text(json.dumps(existing_data, ensure_ascii=False, indent=2), encoding="utf-8")
path.parent.mkdir(parents=True, exist_ok=True)
data = _load_data()
if any(key in data for key in ("default_model", "current_sessions", "session_models")):
data = {"$default": LLMState.from_dict(data).to_dict()}
data[channel] = state.to_dict()
path.write_text(json.dumps(data, ensure_ascii=False, indent=2), encoding="utf-8")


def get_default_model(channel: str = "$default") -> str | None:
Expand All @@ -64,3 +98,35 @@ def set_default_model(model_name: str | None, channel: str = "$default") -> None
state = _read_state(channel)
state.default_model = model_name if model_name else None
_write_state(state, channel)


def get_current_session(user_id: str) -> SessionPointer | None:
return _read_state("$default").current_sessions.get(user_id)


def set_current_session(user_id: str, session_id: str, created_at: int) -> None:
state = _read_state("$default")
state.current_sessions[user_id] = SessionPointer(session_id, created_at)
_write_state(state, "$default")


def clear_current_session(user_id: str) -> None:
state = _read_state("$default")
if state.current_sessions.pop(user_id, None) is not None:
_write_state(state, "$default")


def get_session_model(session_id: str) -> str | None:
return _read_state("$default").session_models.get(session_id)


def set_session_model(session_id: str, model: str) -> None:
state = _read_state("$default")
state.session_models[session_id] = model
_write_state(state, "$default")


def clear_session_model(session_id: str) -> None:
state = _read_state("$default")
if state.session_models.pop(session_id, None) is not None:
_write_state(state, "$default")
35 changes: 0 additions & 35 deletions src/entari_plugin_llm/_types.py

This file was deleted.

4 changes: 0 additions & 4 deletions src/entari_plugin_llm/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -32,10 +32,6 @@ class Config(BasicConfModel, extra="allow"):
"""全局提示词。用作没有特定提示词的模型的后备"""
models: list[ScopedModel] = model_field(default_factory=list)
"""配置模型及其各自设置的列表"""
toolcall_max_steps: int = 8
"""单个会话中工具调用的最大步骤数"""
context_length: int = 50
"""上下文长度"""
tools: dict[str, dict[str, Any]] = model_field(default_factory=dict)
"""工具"""

Expand Down
7 changes: 4 additions & 3 deletions src/entari_plugin_llm/event.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,24 +5,25 @@
from arclet.entari.const import ITEM_ACCOUNT, ITEM_SESSION
from arclet.letoderea import Contexts, Result, define, provide

from .model import LLMSession
from .sessions import SessionInfo


@dataclass
class LLMCollectVariableEvent:
session: Session
llm_session: LLMSession
llm_session: SessionInfo
user_message: MessageChain

def check_result(self, value) -> Result[dict[str, Any]] | None:
if isinstance(value, dict):
return Result(value)
return None


collect_vars = define(LLMCollectVariableEvent, name="llm/collect_vars")
collect_vars.providers.extend(
[
provide(LLMSession, call="$llm_session"),
provide(SessionInfo, call="$llm_session"),
provide(MessageChain, call="$user_message"),
]
)
Expand Down
3 changes: 1 addition & 2 deletions src/entari_plugin_llm/handlers/chat.py
Original file line number Diff line number Diff line change
Expand Up @@ -60,5 +60,4 @@ async def reload_config(event: ConfigReload):
new_conf = config_model_validate(Config, event.value)
_conf.models = new_conf.models
_conf.prompt = new_conf.prompt
_conf.context_length = new_conf.context_length
_conf.toolcall_max_steps = new_conf.toolcall_max_steps
_conf.tools = new_conf.tools
56 changes: 0 additions & 56 deletions src/entari_plugin_llm/json_output.py

This file was deleted.

Loading