Skip to main content

systemprompt_agent/
extension.rs

1//! `Extension` registration for the agent domain: schemas, migrations,
2//! dependencies.
3//!
4//! Copyright (c) systemprompt.io — Business Source License 1.1.
5//! See <https://systemprompt.io> for licensing details.
6
7use systemprompt_extension::prelude::*;
8
9#[derive(Debug, Clone, Copy, Default)]
10pub struct AgentExtension;
11
12impl Extension for AgentExtension {
13    fn metadata(&self) -> ExtensionMetadata {
14        ExtensionMetadata {
15            id: "agent",
16            name: "Agent",
17            version: env!("CARGO_PKG_VERSION"),
18        }
19    }
20
21    fn schemas(&self) -> Vec<SchemaDefinition> {
22        let mut schemas = vec![SchemaDefinition::sql_only(include_str!(
23            "../schema/reporting_capture.sql"
24        ))];
25        schemas.extend(conversation_schemas());
26        schemas.extend(artifact_schemas());
27        schemas.extend(context_schemas());
28        schemas.extend(task_tracking_schemas());
29        schemas
30    }
31
32    fn dependencies(&self) -> Vec<&'static str> {
33        vec!["users", "oauth", "mcp", "ai"]
34    }
35
36    fn migrations(&self) -> Vec<Migration> {
37        extension_migrations!()
38    }
39
40    fn cross_extension_tables(&self) -> Vec<&'static str> {
41        vec!["ai_requests", "services"]
42    }
43}
44
45fn conversation_schemas() -> Vec<SchemaDefinition> {
46    vec![
47        SchemaDefinition::sql_only(include_str!("../schema/reporting_privacy.sql")),
48        SchemaDefinition::new("user_contexts", include_str!("../schema/user_contexts.sql"))
49            .with_required_columns(vec![
50                "context_id".into(),
51                "user_id".into(),
52                "created_at".into(),
53            ]),
54        SchemaDefinition::new("agent_tasks", include_str!("../schema/agent_tasks.sql"))
55            .with_required_columns(vec![
56                "task_id".into(),
57                "context_id".into(),
58                "status".into(),
59                "created_at".into(),
60            ]),
61        SchemaDefinition::new("task_messages", include_str!("../schema/task_messages.sql"))
62            .with_required_columns(vec![
63                "id".into(),
64                "task_id".into(),
65                "role".into(),
66                "created_at".into(),
67            ]),
68        SchemaDefinition::new("message_parts", include_str!("../schema/message_parts.sql"))
69            .with_required_columns(vec!["id".into(), "message_id".into(), "part_kind".into()]),
70    ]
71}
72
73fn context_schemas() -> Vec<SchemaDefinition> {
74    vec![
75        SchemaDefinition::new(
76            "context_agents",
77            include_str!("../schema/context_agents.sql"),
78        )
79        .with_required_columns(vec!["id".into(), "context_id".into(), "agent_name".into()]),
80        SchemaDefinition::new(
81            "context_notifications",
82            include_str!("../schema/context_notifications.sql"),
83        )
84        .with_required_columns(vec![
85            "id".into(),
86            "context_id".into(),
87            "notification_type".into(),
88        ]),
89    ]
90}
91
92fn task_tracking_schemas() -> Vec<SchemaDefinition> {
93    vec![
94        SchemaDefinition::new(
95            "task_execution_steps",
96            include_str!("../schema/task_execution_steps.sql"),
97        )
98        .with_required_columns(vec![
99            "step_id".into(),
100            "task_id".into(),
101            "step_type".into(),
102        ]),
103    ]
104}
105
106fn artifact_schemas() -> Vec<SchemaDefinition> {
107    vec![
108        SchemaDefinition::new(
109            "task_artifacts",
110            include_str!("../schema/task_artifacts.sql"),
111        )
112        .with_required_columns(vec!["id".into(), "task_id".into(), "artifact_id".into()]),
113        SchemaDefinition::new(
114            "artifact_parts",
115            include_str!("../schema/artifact_parts.sql"),
116        )
117        .with_required_columns(vec!["id".into(), "artifact_id".into(), "part_kind".into()]),
118    ]
119}
120
121register_extension!(AgentExtension);