feat(campaign): support per-campaign analysis model override
Campaigns can pin an analysis model config instead of following the global analysis default. Creation validates the referenced config exists (400 otherwise); the create form offers enabled chat configs with the global default as the fallback option.
This commit is contained in:
parent
e1e067bac4
commit
c63236f671
@ -289,6 +289,7 @@ class Campaign(BaseModel):
|
|||||||
completed_at: Optional[datetime] = None
|
completed_at: Optional[datetime] = None
|
||||||
created_at: Optional[datetime] = None
|
created_at: Optional[datetime] = None
|
||||||
summary: Optional[CampaignSummary] = None
|
summary: Optional[CampaignSummary] = None
|
||||||
|
analysis_model_config_id: Optional[str] = None
|
||||||
|
|
||||||
|
|
||||||
class Turn(BaseModel):
|
class Turn(BaseModel):
|
||||||
|
|||||||
@ -187,6 +187,7 @@ class CampaignDB(SQLModel, table=True):
|
|||||||
completed_at: Optional[datetime] = None
|
completed_at: Optional[datetime] = None
|
||||||
created_at: Optional[datetime] = Field(default_factory=utc_now)
|
created_at: Optional[datetime] = Field(default_factory=utc_now)
|
||||||
summary: Optional[str] = None
|
summary: Optional[str] = None
|
||||||
|
analysis_model_config_id: Optional[str] = None
|
||||||
|
|
||||||
def get_plan(self) -> list[dict[str, Any]]:
|
def get_plan(self) -> list[dict[str, Any]]:
|
||||||
return _json_loads(self.plan)
|
return _json_loads(self.plan)
|
||||||
|
|||||||
@ -350,6 +350,7 @@ class CampaignRepository(BaseRepository[Campaign, CampaignDB]):
|
|||||||
started_at=campaign.started_at,
|
started_at=campaign.started_at,
|
||||||
completed_at=campaign.completed_at,
|
completed_at=campaign.completed_at,
|
||||||
created_at=campaign.created_at,
|
created_at=campaign.created_at,
|
||||||
|
analysis_model_config_id=campaign.analysis_model_config_id,
|
||||||
)
|
)
|
||||||
db.set_plan([entry.model_dump(mode="json") for entry in campaign.plan])
|
db.set_plan([entry.model_dump(mode="json") for entry in campaign.plan])
|
||||||
if campaign.summary:
|
if campaign.summary:
|
||||||
@ -369,6 +370,7 @@ class CampaignRepository(BaseRepository[Campaign, CampaignDB]):
|
|||||||
completed_at=db.completed_at,
|
completed_at=db.completed_at,
|
||||||
created_at=db.created_at,
|
created_at=db.created_at,
|
||||||
summary=db.get_summary(),
|
summary=db.get_summary(),
|
||||||
|
analysis_model_config_id=db.analysis_model_config_id,
|
||||||
)
|
)
|
||||||
|
|
||||||
def update(self, campaign: Campaign) -> Optional[Campaign]:
|
def update(self, campaign: Campaign) -> Optional[Campaign]:
|
||||||
@ -383,6 +385,7 @@ class CampaignRepository(BaseRepository[Campaign, CampaignDB]):
|
|||||||
existing.status = campaign.status.value
|
existing.status = campaign.status.value
|
||||||
existing.started_at = campaign.started_at
|
existing.started_at = campaign.started_at
|
||||||
existing.completed_at = campaign.completed_at
|
existing.completed_at = campaign.completed_at
|
||||||
|
existing.analysis_model_config_id = campaign.analysis_model_config_id
|
||||||
if campaign.summary is not None:
|
if campaign.summary is not None:
|
||||||
existing.set_summary(campaign.summary.model_dump(mode="json"))
|
existing.set_summary(campaign.summary.model_dump(mode="json"))
|
||||||
self.session.add(existing)
|
self.session.add(existing)
|
||||||
|
|||||||
@ -20,6 +20,7 @@ from agenteval.evaluation.report import (
|
|||||||
from agenteval.evaluation.report_render import render_campaign_markdown
|
from agenteval.evaluation.report_render import render_campaign_markdown
|
||||||
from agenteval.models import Campaign, CampaignPlanEntry, CampaignStatus
|
from agenteval.models import Campaign, CampaignPlanEntry, CampaignStatus
|
||||||
from agenteval.storage.db import utc_now
|
from agenteval.storage.db import utc_now
|
||||||
|
from agenteval.storage.model_config_repository import ModelConfigRepository
|
||||||
from agenteval.storage.repository import (
|
from agenteval.storage.repository import (
|
||||||
CampaignRepository,
|
CampaignRepository,
|
||||||
RunRepository,
|
RunRepository,
|
||||||
@ -37,6 +38,7 @@ class CreateCampaignRequest(BaseModel):
|
|||||||
window_seconds: int = Field(gt=0)
|
window_seconds: int = Field(gt=0)
|
||||||
time_scale: float = Field(default=1.0, gt=0)
|
time_scale: float = Field(default=1.0, gt=0)
|
||||||
plan: list[CampaignPlanEntry] = Field(min_length=1)
|
plan: list[CampaignPlanEntry] = Field(min_length=1)
|
||||||
|
analysis_model_config_id: str | None = None
|
||||||
|
|
||||||
|
|
||||||
@router.get("")
|
@router.get("")
|
||||||
@ -68,12 +70,17 @@ async def create_campaign(
|
|||||||
detail=f"scenario not found: {entry.scenario_id}",
|
detail=f"scenario not found: {entry.scenario_id}",
|
||||||
)
|
)
|
||||||
|
|
||||||
|
if request.analysis_model_config_id is not None:
|
||||||
|
if not ModelConfigRepository(session).get(request.analysis_model_config_id):
|
||||||
|
raise HTTPException(status_code=400, detail="analysis model config not found")
|
||||||
|
|
||||||
campaign = Campaign(
|
campaign = Campaign(
|
||||||
name=request.name,
|
name=request.name,
|
||||||
target_id=request.target_id,
|
target_id=request.target_id,
|
||||||
window_seconds=request.window_seconds,
|
window_seconds=request.window_seconds,
|
||||||
time_scale=request.time_scale,
|
time_scale=request.time_scale,
|
||||||
plan=request.plan,
|
plan=request.plan,
|
||||||
|
analysis_model_config_id=request.analysis_model_config_id,
|
||||||
)
|
)
|
||||||
repo = CampaignRepository(session)
|
repo = CampaignRepository(session)
|
||||||
campaign = repo.create(campaign)
|
campaign = repo.create(campaign)
|
||||||
|
|||||||
@ -336,6 +336,7 @@ export interface Campaign {
|
|||||||
started_at: string | null
|
started_at: string | null
|
||||||
completed_at: string | null
|
completed_at: string | null
|
||||||
summary: Record<string, unknown> | null
|
summary: Record<string, unknown> | null
|
||||||
|
analysis_model_config_id: string | null
|
||||||
progress?: CampaignProgress
|
progress?: CampaignProgress
|
||||||
}
|
}
|
||||||
|
|
||||||
@ -399,6 +400,7 @@ export interface CreateCampaignPayload {
|
|||||||
window_seconds: number
|
window_seconds: number
|
||||||
time_scale: number
|
time_scale: number
|
||||||
plan: CampaignPlanEntry[]
|
plan: CampaignPlanEntry[]
|
||||||
|
analysis_model_config_id?: string | null
|
||||||
}
|
}
|
||||||
|
|
||||||
export const campaignsApi = {
|
export const campaignsApi = {
|
||||||
|
|||||||
@ -13,9 +13,9 @@ import { Line, Bar } from '@ant-design/charts'
|
|||||||
import PageWrapper from '../components/PageWrapper'
|
import PageWrapper from '../components/PageWrapper'
|
||||||
import StatCard from '../components/StatCard'
|
import StatCard from '../components/StatCard'
|
||||||
import {
|
import {
|
||||||
campaignsApi, targetsApi, scenariosApi, runsApi,
|
campaignsApi, targetsApi, scenariosApi, runsApi, modelConfigsApi,
|
||||||
type CampaignListItem, type CampaignReport, type Target, type Scenario, type Run,
|
type CampaignListItem, type CampaignReport, type Target, type Scenario, type Run,
|
||||||
type CampaignTimelineEntry,
|
type CampaignTimelineEntry, type ModelConfig,
|
||||||
} from '../api'
|
} from '../api'
|
||||||
import { passRateColor } from '../utils/colors'
|
import { passRateColor } from '../utils/colors'
|
||||||
import { shortDateTime, toDate } from '../utils/date'
|
import { shortDateTime, toDate } from '../utils/date'
|
||||||
@ -94,6 +94,7 @@ interface CampaignsListsData {
|
|||||||
campaigns: CampaignListItem[]
|
campaigns: CampaignListItem[]
|
||||||
targets: Target[]
|
targets: Target[]
|
||||||
scenarios: Scenario[]
|
scenarios: Scenario[]
|
||||||
|
modelConfigs: ModelConfig[]
|
||||||
}
|
}
|
||||||
|
|
||||||
export default function CampaignsPage() {
|
export default function CampaignsPage() {
|
||||||
@ -123,14 +124,18 @@ export default function CampaignsPage() {
|
|||||||
|
|
||||||
const { data, loading, reload } = useResource<CampaignsListsData>(
|
const { data, loading, reload } = useResource<CampaignsListsData>(
|
||||||
async () => {
|
async () => {
|
||||||
const [c, t, s] = await Promise.all([campaignsApi.list(), targetsApi.list(), scenariosApi.list()])
|
const [c, t, s, m] = await Promise.all([
|
||||||
return { campaigns: c.data, targets: t.data, scenarios: s.data }
|
campaignsApi.list(), targetsApi.list(), scenariosApi.list(), modelConfigsApi.list(),
|
||||||
|
])
|
||||||
|
return { campaigns: c.data, targets: t.data, scenarios: s.data, modelConfigs: m.data }
|
||||||
},
|
},
|
||||||
{ tabPath: '/campaigns' },
|
{ tabPath: '/campaigns' },
|
||||||
)
|
)
|
||||||
const campaigns = data?.campaigns ?? []
|
const campaigns = data?.campaigns ?? []
|
||||||
const targets = data?.targets ?? []
|
const targets = data?.targets ?? []
|
||||||
const scenarios = data?.scenarios ?? []
|
const scenarios = data?.scenarios ?? []
|
||||||
|
const chatConfigs = (data?.modelConfigs ?? []).filter((m) => m.enabled && m.capability === 'chat')
|
||||||
|
const analysisDefault = chatConfigs.find((m) => m.is_analysis_default)
|
||||||
const scenarioNames = Object.fromEntries(scenarios.map((s) => [s.id, s.name]))
|
const scenarioNames = Object.fromEntries(scenarios.map((s) => [s.id, s.name]))
|
||||||
|
|
||||||
const wPlan = (Form.useWatch('plan', form) as PlanFormEntry[] | undefined) ?? []
|
const wPlan = (Form.useWatch('plan', form) as PlanFormEntry[] | undefined) ?? []
|
||||||
@ -163,6 +168,7 @@ export default function CampaignsPage() {
|
|||||||
name: '', target_id: undefined, window_seconds: 24 * 3600,
|
name: '', target_id: undefined, window_seconds: 24 * 3600,
|
||||||
realtime: false, target_value: 60, target_unit: 60,
|
realtime: false, target_value: 60, target_unit: 60,
|
||||||
plan: [{ scenario_id: undefined, offset_hours: 0, count: 1 }],
|
plan: [{ scenario_id: undefined, offset_hours: 0, count: 1 }],
|
||||||
|
analysis_model_config_id: null,
|
||||||
})
|
})
|
||||||
setCreateOpen(true)
|
setCreateOpen(true)
|
||||||
}
|
}
|
||||||
@ -185,6 +191,7 @@ export default function CampaignsPage() {
|
|||||||
offset_seconds: Math.round((e.offset_hours ?? 0) * 3600),
|
offset_seconds: Math.round((e.offset_hours ?? 0) * 3600),
|
||||||
count: e.count ?? 1,
|
count: e.count ?? 1,
|
||||||
})),
|
})),
|
||||||
|
analysis_model_config_id: (values.analysis_model_config_id as string | null) ?? null,
|
||||||
})
|
})
|
||||||
message.success('评估活动已创建并开始调度')
|
message.success('评估活动已创建并开始调度')
|
||||||
setCreateOpen(false)
|
setCreateOpen(false)
|
||||||
@ -576,6 +583,26 @@ export default function CampaignsPage() {
|
|||||||
</>
|
</>
|
||||||
)}
|
)}
|
||||||
</div>
|
</div>
|
||||||
|
<Form.Item
|
||||||
|
name="analysis_model_config_id"
|
||||||
|
label="分析模型"
|
||||||
|
tooltip="活动完成后生成智能分析报告使用的模型;不选则跟随模型配置中心的「分析默认」"
|
||||||
|
style={{ marginTop: 16, marginBottom: 0 }}
|
||||||
|
>
|
||||||
|
<Select
|
||||||
|
showSearch
|
||||||
|
optionFilterProp="label"
|
||||||
|
options={[
|
||||||
|
{
|
||||||
|
label: analysisDefault
|
||||||
|
? `全局默认(${analysisDefault.name})`
|
||||||
|
: '全局默认(未配置)',
|
||||||
|
value: null,
|
||||||
|
},
|
||||||
|
...chatConfigs.map((m) => ({ label: m.name, value: m.id as string | null })),
|
||||||
|
]}
|
||||||
|
/>
|
||||||
|
</Form.Item>
|
||||||
</Col>
|
</Col>
|
||||||
|
|
||||||
<Col span={15}>
|
<Col span={15}>
|
||||||
|
|||||||
@ -0,0 +1,29 @@
|
|||||||
|
"""add analysis model override to campaigns
|
||||||
|
|
||||||
|
Revision ID: b7e4d3a92c15
|
||||||
|
Revises: a3f9c2e71b48
|
||||||
|
Create Date: 2026-08-03
|
||||||
|
"""
|
||||||
|
|
||||||
|
from typing import Sequence, Union
|
||||||
|
|
||||||
|
import sqlalchemy as sa
|
||||||
|
import sqlmodel # noqa: F401
|
||||||
|
from alembic import op
|
||||||
|
|
||||||
|
revision: str = "b7e4d3a92c15"
|
||||||
|
down_revision: Union[str, Sequence[str], None] = "a3f9c2e71b48"
|
||||||
|
branch_labels: Union[str, Sequence[str], None] = None
|
||||||
|
depends_on: Union[str, Sequence[str], None] = None
|
||||||
|
|
||||||
|
|
||||||
|
def upgrade() -> None:
|
||||||
|
with op.batch_alter_table("campaigns", schema=None) as batch_op:
|
||||||
|
batch_op.add_column(
|
||||||
|
sa.Column("analysis_model_config_id", sqlmodel.sql.sqltypes.AutoString(), nullable=True)
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def downgrade() -> None:
|
||||||
|
with op.batch_alter_table("campaigns", schema=None) as batch_op:
|
||||||
|
batch_op.drop_column("analysis_model_config_id")
|
||||||
@ -332,3 +332,43 @@ async def test_campaign_timeline_structure(client, seeded_db):
|
|||||||
async def test_campaign_timeline_missing_404(client, seeded_db):
|
async def test_campaign_timeline_missing_404(client, seeded_db):
|
||||||
resp = await client.get("/api/campaigns/nope/timeline")
|
resp = await client.get("/api/campaigns/nope/timeline")
|
||||||
assert resp.status_code == 404
|
assert resp.status_code == 404
|
||||||
|
|
||||||
|
|
||||||
|
# ── analysis model override (v0.7 ticket 02) ─────────────────────────────
|
||||||
|
|
||||||
|
def _seed_chat_config(session, config_id: str = "mc-1") -> None:
|
||||||
|
from agenteval.storage.db import ModelConfigDB
|
||||||
|
from agenteval.storage.model_config_repository import ModelConfigRepository
|
||||||
|
|
||||||
|
ModelConfigRepository(session).create(ModelConfigDB(
|
||||||
|
id=config_id, name="chat-cfg", provider="openai_compatible", capability="chat",
|
||||||
|
endpoint_url="https://models.example.com/v1/chat/completions", model_name="chat-model",
|
||||||
|
))
|
||||||
|
|
||||||
|
|
||||||
|
async def test_create_campaign_with_analysis_model_override(client, seeded_db):
|
||||||
|
_seed_chat_config(seeded_db)
|
||||||
|
resp = await client.post("/api/campaigns", json=_valid_payload(analysis_model_config_id="mc-1"))
|
||||||
|
assert resp.status_code == 200
|
||||||
|
assert resp.json()["analysis_model_config_id"] == "mc-1"
|
||||||
|
|
||||||
|
campaign_id = resp.json()["id"]
|
||||||
|
got = (await client.get(f"/api/campaigns/{campaign_id}")).json()
|
||||||
|
assert got["analysis_model_config_id"] == "mc-1"
|
||||||
|
listing = (await client.get("/api/campaigns")).json()
|
||||||
|
assert listing[0]["analysis_model_config_id"] == "mc-1"
|
||||||
|
|
||||||
|
|
||||||
|
async def test_create_campaign_without_override_stores_null(client, seeded_db):
|
||||||
|
resp = await client.post("/api/campaigns", json=_valid_payload())
|
||||||
|
assert resp.status_code == 200
|
||||||
|
assert resp.json()["analysis_model_config_id"] is None
|
||||||
|
|
||||||
|
campaign_id = resp.json()["id"]
|
||||||
|
got = (await client.get(f"/api/campaigns/{campaign_id}")).json()
|
||||||
|
assert got["analysis_model_config_id"] is None
|
||||||
|
|
||||||
|
|
||||||
|
async def test_create_campaign_invalid_analysis_model_400(client, seeded_db):
|
||||||
|
resp = await client.post("/api/campaigns", json=_valid_payload(analysis_model_config_id="nope"))
|
||||||
|
assert resp.status_code == 400
|
||||||
|
|||||||
38
tests/unit/test_campaign_analysis_model_migration.py
Normal file
38
tests/unit/test_campaign_analysis_model_migration.py
Normal file
@ -0,0 +1,38 @@
|
|||||||
|
"""Verify the campaign analysis-model migration adds a nullable column."""
|
||||||
|
|
||||||
|
import importlib
|
||||||
|
|
||||||
|
import sqlalchemy as sa
|
||||||
|
from alembic.migration import MigrationContext
|
||||||
|
from alembic.operations import Operations
|
||||||
|
|
||||||
|
|
||||||
|
def test_campaign_analysis_model_migration_adds_nullable_column(tmp_path, monkeypatch):
|
||||||
|
engine = sa.create_engine(f"sqlite:///{tmp_path / 'campaign_analysis.db'}")
|
||||||
|
metadata = sa.MetaData()
|
||||||
|
legacy_table = sa.Table(
|
||||||
|
"campaigns",
|
||||||
|
metadata,
|
||||||
|
sa.Column("id", sa.String(), primary_key=True),
|
||||||
|
sa.Column("name", sa.String(), nullable=False),
|
||||||
|
)
|
||||||
|
metadata.create_all(engine)
|
||||||
|
|
||||||
|
with engine.begin() as connection:
|
||||||
|
connection.execute(legacy_table.insert().values(id="camp-1", name="现有活动"))
|
||||||
|
operations = Operations(MigrationContext.configure(connection))
|
||||||
|
migration = importlib.import_module(
|
||||||
|
"migrations.versions.b7e4d3a92c15_add_campaign_analysis_model"
|
||||||
|
)
|
||||||
|
monkeypatch.setattr(migration, "op", operations)
|
||||||
|
migration.upgrade()
|
||||||
|
|
||||||
|
columns = {c["name"]: c for c in sa.inspect(connection).get_columns("campaigns")}
|
||||||
|
assert "analysis_model_config_id" in columns
|
||||||
|
assert columns["analysis_model_config_id"]["nullable"] is True
|
||||||
|
|
||||||
|
value = connection.execute(
|
||||||
|
sa.text("SELECT analysis_model_config_id FROM campaigns WHERE id = :id"),
|
||||||
|
{"id": "camp-1"},
|
||||||
|
).scalar_one()
|
||||||
|
assert value is None
|
||||||
Loading…
Reference in New Issue
Block a user