Files
agent_management/agent_templates/agents/coding_a2a_agent/config.py
T

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, "")}