347 lines
12 KiB
Python
347 lines
12 KiB
Python
"""
|
|
Configuration models for the coding A2A agent.
|
|
"""
|
|
from __future__ import annotations
|
|
|
|
import os
|
|
from pathlib import Path
|
|
from enum import Enum
|
|
from typing import Any, Optional
|
|
|
|
from pydantic import BaseModel, Field, model_validator
|
|
|
|
|
|
class GitProvider(str, Enum):
|
|
gitea = "gitea"
|
|
github = "github"
|
|
gitlab = "gitlab"
|
|
generic = "generic"
|
|
|
|
|
|
class DatabaseEngine(str, Enum):
|
|
mysql = "mysql"
|
|
postgresql = "postgresql"
|
|
|
|
|
|
def _env_text(*names: str) -> Optional[str]:
|
|
for name in names:
|
|
value = os.getenv(name)
|
|
if value is not None and value != "":
|
|
return value
|
|
return None
|
|
|
|
|
|
def _env_int(*names: str) -> Optional[int]:
|
|
value = _env_text(*names)
|
|
return int(value) if value is not None else None
|
|
|
|
|
|
def _env_list(*names: str) -> list[str]:
|
|
value = _env_text(*names)
|
|
if not value:
|
|
return []
|
|
return [item.strip() for item in value.split(",") if item.strip()]
|
|
|
|
|
|
class LiteLLMConfig(BaseModel):
|
|
base_url: str = Field(
|
|
default_factory=lambda: (
|
|
os.getenv("LITELLM_BASE_URL")
|
|
or os.getenv("LLM_BASE_URL")
|
|
or os.getenv("OPENAI_BASE_URL")
|
|
or "https://litellm.graystone-fb459c5d.southeastasia.azurecontainerapps.io/v1"
|
|
).rstrip("/")
|
|
)
|
|
api_key: Optional[str] = Field(default_factory=lambda: os.getenv("LITELLM_API_KEY") or os.getenv("OPENAI_API_KEY"))
|
|
model: str = Field(
|
|
default_factory=lambda: (
|
|
os.getenv("MODEL_NAME")
|
|
or os.getenv("LITELLM_MODEL")
|
|
or os.getenv("LLM_MODEL")
|
|
or "taiji/gpt-4o-mini"
|
|
)
|
|
)
|
|
timeout: int = Field(default_factory=lambda: int(os.getenv("LITELLM_TIMEOUT") or os.getenv("LLM_TIMEOUT") or "600"))
|
|
max_tokens: int = Field(default_factory=lambda: int(os.getenv("LITELLM_MAX_TOKENS") or os.getenv("LLM_MAX_TOKENS") or "4096"))
|
|
|
|
@property
|
|
def normalized_model(self) -> str:
|
|
if ":" in self.model:
|
|
return self.model
|
|
return f"openai:{self.model}"
|
|
|
|
|
|
class WorkspaceConfig(BaseModel):
|
|
root_dir: str = Field(default_factory=lambda: os.getenv("WORK_DIR", "/workspace"))
|
|
entry_file: Optional[str] = None
|
|
context_files: list[str] = Field(default_factory=list)
|
|
allowed_paths: list[str] = Field(default_factory=list)
|
|
|
|
|
|
class GitResourceConfig(BaseModel):
|
|
provider: Optional[GitProvider] = None
|
|
repo_url: Optional[str] = None
|
|
default_branch: str = "main"
|
|
username: Optional[str] = None
|
|
password: Optional[str] = None
|
|
token: Optional[str] = None
|
|
local_path: Optional[str] = None
|
|
allowed_paths: list[str] = Field(default_factory=list)
|
|
write_mode: str = "branch"
|
|
|
|
@model_validator(mode="after")
|
|
def infer_provider(self) -> "GitResourceConfig":
|
|
if self.provider is None and self.repo_url:
|
|
lowered = self.repo_url.lower()
|
|
if "github" in lowered:
|
|
self.provider = GitProvider.github
|
|
elif "gitlab" in lowered:
|
|
self.provider = GitProvider.gitlab
|
|
elif "gitea" in lowered or ":3000/" in lowered or "/api/v1/" in lowered:
|
|
self.provider = GitProvider.gitea
|
|
else:
|
|
self.provider = GitProvider.generic
|
|
return self
|
|
|
|
|
|
class DatabaseResourceConfig(BaseModel):
|
|
engine: DatabaseEngine
|
|
host: str
|
|
port: Optional[int] = None
|
|
username: str
|
|
password: str
|
|
database: str
|
|
ssl_mode: Optional[str] = None
|
|
|
|
|
|
class AzureBlobResourceConfig(BaseModel):
|
|
account_url: Optional[str] = None
|
|
connection_string: Optional[str] = None
|
|
container_name: str
|
|
account_name: Optional[str] = None
|
|
account_key: Optional[str] = None
|
|
sas_token: Optional[str] = None
|
|
prefix: str = ""
|
|
|
|
|
|
class ResourceConfig(BaseModel):
|
|
git: Optional[GitResourceConfig] = None
|
|
mysql: Optional[DatabaseResourceConfig] = None
|
|
postgresql: Optional[DatabaseResourceConfig] = None
|
|
azure_blob: Optional[AzureBlobResourceConfig] = None
|
|
|
|
@model_validator(mode="before")
|
|
@classmethod
|
|
def apply_env_defaults(cls, data: Any) -> Any:
|
|
if isinstance(data, cls):
|
|
return data
|
|
|
|
payload = dict(data or {})
|
|
git_env = _git_resource_from_env()
|
|
mysql_env = _mysql_resource_from_env()
|
|
postgres_env = _postgres_resource_from_env()
|
|
blob_env = _azure_blob_resource_from_env()
|
|
|
|
if "git" not in payload and git_env:
|
|
payload["git"] = git_env
|
|
elif isinstance(payload.get("git"), dict) and git_env:
|
|
payload["git"] = {**git_env, **payload["git"]}
|
|
|
|
if "mysql" not in payload and mysql_env:
|
|
payload["mysql"] = mysql_env
|
|
elif isinstance(payload.get("mysql"), dict) and mysql_env:
|
|
payload["mysql"] = {**mysql_env, **payload["mysql"]}
|
|
|
|
if "postgresql" not in payload and postgres_env:
|
|
payload["postgresql"] = postgres_env
|
|
elif isinstance(payload.get("postgresql"), dict) and postgres_env:
|
|
payload["postgresql"] = {**postgres_env, **payload["postgresql"]}
|
|
|
|
if "azure_blob" not in payload and blob_env:
|
|
payload["azure_blob"] = blob_env
|
|
elif isinstance(payload.get("azure_blob"), dict) and blob_env:
|
|
payload["azure_blob"] = {**blob_env, **payload["azure_blob"]}
|
|
|
|
return payload
|
|
|
|
@property
|
|
def enabled_resource_names(self) -> list[str]:
|
|
names: list[str] = []
|
|
if self.git:
|
|
names.append("git")
|
|
if self.mysql:
|
|
names.append("mysql")
|
|
if self.postgresql:
|
|
names.append("postgresql")
|
|
if self.azure_blob:
|
|
names.append("azure_blob")
|
|
return names
|
|
|
|
|
|
class AgentMetadata(BaseModel):
|
|
name: str = Field(default_factory=lambda: os.getenv("AGENT_NAME", "coding-a2a-agent"))
|
|
description: str = Field(
|
|
default=(
|
|
"Claude Code 风格的编程 Agent,使用 Pydantic AI 作为核心,"
|
|
"支持 A2A 协议,以及 Git / DB / Azure Blob 资源工具。"
|
|
)
|
|
)
|
|
version: str = "1.0.0"
|
|
enable_streaming: bool = True
|
|
role_name: Optional[str] = Field(
|
|
default_factory=lambda: os.getenv("AGENT_ROLE_NAME") or os.getenv("AGENT_ROLE")
|
|
)
|
|
instruction_text: Optional[str] = Field(
|
|
default_factory=lambda: os.getenv("AGENT_INSTRUCTION_TEXT")
|
|
)
|
|
instruction_file: Optional[str] = Field(
|
|
default_factory=lambda: os.getenv("AGENT_INSTRUCTION_FILE")
|
|
)
|
|
system_prompt: str = Field(
|
|
default=(
|
|
"You are a senior coding agent similar to Claude Code. "
|
|
"Understand the repository first, then make minimal precise changes. "
|
|
"Prefer using tools to inspect, edit, run checks, inspect git state, "
|
|
"query configured databases, and inspect Azure Blob artifacts. "
|
|
"Always end by calling finish(summary)."
|
|
)
|
|
)
|
|
instruction_source: str = "default"
|
|
instruction_content: Optional[str] = None
|
|
|
|
@model_validator(mode="after")
|
|
def load_instruction_content(self) -> "AgentMetadata":
|
|
if self.instruction_text and self.instruction_text.strip():
|
|
self.instruction_source = "env_text"
|
|
self.instruction_content = self.instruction_text.strip()
|
|
return self
|
|
|
|
if self.instruction_file:
|
|
instruction_path = Path(self.instruction_file)
|
|
if instruction_path.exists() and instruction_path.is_file():
|
|
self.instruction_source = f"env_file:{instruction_path}"
|
|
self.instruction_content = instruction_path.read_text(
|
|
encoding="utf-8",
|
|
errors="replace",
|
|
).strip()
|
|
return self
|
|
|
|
self.instruction_source = "default"
|
|
self.instruction_content = None
|
|
return self
|
|
|
|
@property
|
|
def effective_system_prompt(self) -> str:
|
|
sections = [self.system_prompt.strip()]
|
|
if self.role_name:
|
|
sections.append(f"Runtime role assignment: {self.role_name.strip()}")
|
|
if self.instruction_content:
|
|
sections.append(
|
|
"Startup instructions loaded from runtime configuration:\n"
|
|
f"{self.instruction_content.strip()}"
|
|
)
|
|
return "\n\n".join(part for part in sections if part)
|
|
|
|
|
|
class CodingRequestConfig(BaseModel):
|
|
workspace: WorkspaceConfig = Field(default_factory=WorkspaceConfig)
|
|
resources: ResourceConfig = Field(default_factory=ResourceConfig)
|
|
task_mode: str = "code"
|
|
branch_name: Optional[str] = None
|
|
commit_message: Optional[str] = None
|
|
model: Optional[str] = None
|
|
api_key: Optional[str] = None
|
|
metadata: dict[str, Any] = Field(default_factory=dict)
|
|
|
|
|
|
def get_runtime_defaults(
|
|
api_key: Optional[str] = None,
|
|
model: Optional[str] = None,
|
|
) -> tuple[LiteLLMConfig, AgentMetadata]:
|
|
llm = LiteLLMConfig(api_key=api_key, model=model or LiteLLMConfig().model)
|
|
meta = AgentMetadata()
|
|
return llm, meta
|
|
|
|
|
|
def _git_resource_from_env() -> Optional[dict[str, Any]]:
|
|
repo_url = _env_text("GIT_REPO_URL")
|
|
username = _env_text("GIT_USERNAME", "GIT_USER")
|
|
password = _env_text("GIT_PASSWORD")
|
|
token = _env_text("GIT_TOKEN", "GITHUB_TOKEN", "GITLAB_TOKEN", "GITEA_TOKEN")
|
|
provider = _env_text("GIT_PROVIDER")
|
|
if not any([repo_url, username, password, token]):
|
|
return None
|
|
data: dict[str, Any] = {
|
|
"repo_url": repo_url,
|
|
"username": username,
|
|
"password": password,
|
|
"token": token,
|
|
"provider": provider,
|
|
"default_branch": _env_text("GIT_DEFAULT_BRANCH") or "main",
|
|
"local_path": _env_text("GIT_LOCAL_PATH"),
|
|
"allowed_paths": _env_list("GIT_ALLOWED_PATHS"),
|
|
"write_mode": _env_text("GIT_WRITE_MODE") or "branch",
|
|
}
|
|
return {key: value for key, value in data.items() if value not in (None, [], "")}
|
|
|
|
|
|
def _mysql_resource_from_env() -> Optional[dict[str, Any]]:
|
|
host = _env_text("MYSQL_HOST")
|
|
username = _env_text("MYSQL_USER", "MYSQL_USERNAME")
|
|
password = _env_text("MYSQL_PASSWORD")
|
|
database = _env_text("MYSQL_DATABASE", "MYSQL_DB")
|
|
if not all([host, username, password, database]):
|
|
return None
|
|
data: dict[str, Any] = {
|
|
"engine": "mysql",
|
|
"host": host,
|
|
"port": _env_int("MYSQL_PORT"),
|
|
"username": username,
|
|
"password": password,
|
|
"database": database,
|
|
"ssl_mode": _env_text("MYSQL_SSL_MODE"),
|
|
}
|
|
return {key: value for key, value in data.items() if value is not None}
|
|
|
|
|
|
def _postgres_resource_from_env() -> Optional[dict[str, Any]]:
|
|
host = _env_text("POSTGRES_HOST", "POSTGRESQL_HOST")
|
|
username = _env_text("POSTGRES_USER", "POSTGRES_USERNAME", "POSTGRESQL_USER")
|
|
password = _env_text("POSTGRES_PASSWORD", "POSTGRESQL_PASSWORD")
|
|
database = _env_text("POSTGRES_DATABASE", "POSTGRES_DB", "POSTGRESQL_DATABASE")
|
|
if not all([host, username, password, database]):
|
|
return None
|
|
data: dict[str, Any] = {
|
|
"engine": "postgresql",
|
|
"host": host,
|
|
"port": _env_int("POSTGRES_PORT", "POSTGRESQL_PORT"),
|
|
"username": username,
|
|
"password": password,
|
|
"database": database,
|
|
"ssl_mode": _env_text("POSTGRES_SSL_MODE", "POSTGRESQL_SSL_MODE"),
|
|
}
|
|
return {key: value for key, value in data.items() if value is not None}
|
|
|
|
|
|
def _azure_blob_resource_from_env() -> Optional[dict[str, Any]]:
|
|
container_name = _env_text("AZURE_BLOB_CONTAINER", "AZURE_STORAGE_CONTAINER")
|
|
connection_string = _env_text("AZURE_BLOB_CONNECTION_STRING", "AZURE_STORAGE_CONNECTION_STRING")
|
|
account_url = _env_text("AZURE_BLOB_ACCOUNT_URL")
|
|
account_name = _env_text("AZURE_BLOB_ACCOUNT_NAME", "AZURE_STORAGE_ACCOUNT_NAME")
|
|
account_key = _env_text("AZURE_BLOB_ACCOUNT_KEY", "AZURE_STORAGE_ACCOUNT_KEY")
|
|
sas_token = _env_text("AZURE_BLOB_SAS_TOKEN")
|
|
if not container_name:
|
|
return None
|
|
if not any([connection_string, account_url, account_name]):
|
|
return None
|
|
data: dict[str, Any] = {
|
|
"container_name": container_name,
|
|
"connection_string": connection_string,
|
|
"account_url": account_url,
|
|
"account_name": account_name,
|
|
"account_key": account_key,
|
|
"sas_token": sas_token,
|
|
"prefix": _env_text("AZURE_BLOB_PREFIX", "AZURE_STORAGE_PREFIX") or "",
|
|
}
|
|
return {key: value for key, value in data.items() if value not in (None, "")}
|