use crate::context::{load_context, truncate_to_budget, ContextRequest};
use crate::error::PromptError;
use crate::ids::{ProviderId, RunId, SessionId};
use crate::skill::SkillSummary;
use async_trait::async_trait;
use std::sync::Arc;
pub struct PromptContext {
pub session_id: SessionId,
pub run_id: RunId,
pub turn: u32,
pub definition_id: String,
pub skill_catalog: Vec<SkillSummary>,
}
pub struct PromptFragment {
pub content: String,
}
#[async_trait]
pub trait PromptLayer: Send + Sync {
fn id(&self) -> &str;
fn priority(&self) -> i32;
async fn render(&self, context: &PromptContext) -> Result<Option<PromptFragment>, PromptError>;
}
pub struct FrameworkBaseLayer;
#[async_trait]
impl PromptLayer for FrameworkBaseLayer {
fn id(&self) -> &str {
"framework-base"
}
fn priority(&self) -> i32 {
0
}
async fn render(
&self,
_context: &PromptContext,
) -> Result<Option<PromptFragment>, PromptError> {
Ok(None)
}
}
pub struct PromptComposer {
layers: Vec<Arc<dyn PromptLayer>>,
providers: Vec<Arc<dyn crate::context::ContextProvider>>,
}
impl PromptComposer {
pub fn new(
layers: Vec<Arc<dyn PromptLayer>>,
providers: Vec<Arc<dyn crate::context::ContextProvider>>,
) -> Self {
Self { layers, providers }
}
pub async fn render_system_prompt(
&self,
context: &PromptContext,
current_provider: Option<&ProviderId>,
) -> Result<String, PromptError> {
let mut ordered: Vec<&Arc<dyn PromptLayer>> = self.layers.iter().collect();
ordered.sort_by_key(|layer| layer.priority());
let mut sections = Vec::new();
for layer in ordered {
if let Some(fragment) = layer.render(context).await? {
sections.push(fragment.content);
}
}
if !self.providers.is_empty() {
let request = ContextRequest {
session_id: context.session_id.clone(),
run_id: context.run_id.clone(),
turn: context.turn,
};
let items = load_context(&self.providers, &request)
.await
.map_err(|error| PromptError::Render(error.to_string()))?;
for item in items {
if let (Some(allowed), Some(current)) = (&item.allowed_providers, current_provider)
{
if !allowed.contains(current) {
continue;
}
}
let text = truncate_to_budget(&item.text, item.budget_tokens);
sections.push(format!("## {} (v{})\n{}", item.source, item.version, text));
}
}
Ok(sections.join("\n\n"))
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::context::{ContextFailurePolicy, ContextItem};
use crate::error::ContextError;
use async_trait::async_trait;
fn ctx() -> PromptContext {
PromptContext {
session_id: SessionId::from("s"),
run_id: RunId::from("r"),
turn: 1,
definition_id: "defs/agent".into(),
skill_catalog: Vec::new(),
}
}
struct StaticLayer {
id: &'static str,
priority: i32,
content: Option<String>,
}
#[async_trait]
impl PromptLayer for StaticLayer {
fn id(&self) -> &str {
self.id
}
fn priority(&self) -> i32 {
self.priority
}
async fn render(
&self,
_context: &PromptContext,
) -> Result<Option<PromptFragment>, PromptError> {
Ok(self
.content
.clone()
.map(|content| PromptFragment { content }))
}
}
fn layer(id: &'static str, priority: i32, content: Option<&str>) -> Arc<dyn PromptLayer> {
Arc::new(StaticLayer {
id,
priority,
content: content.map(str::to_string),
})
}
struct FixedProvider {
items: Vec<ContextItem>,
}
#[async_trait]
impl crate::context::ContextProvider for FixedProvider {
fn id(&self) -> &str {
"fixed"
}
async fn load(&self, _request: &ContextRequest) -> Result<Vec<ContextItem>, ContextError> {
Ok(self.items.clone())
}
}
fn item(source: &str, version: &str, text: &str) -> ContextItem {
ContextItem {
source: source.into(),
version: version.into(),
text: text.into(),
budget_tokens: None,
persistable: true,
allowed_providers: None,
failure: ContextFailurePolicy::Block,
}
}
#[tokio::test]
async fn layers_render_in_priority_then_registration_order() {
let composer = PromptComposer::new(
vec![
layer("late", 200, Some("late")),
layer("first-a", 100, Some("a")),
layer("first-b", 100, Some("b")),
],
Vec::new(),
);
let prompt = composer.render_system_prompt(&ctx(), None).await.unwrap();
assert_eq!(prompt, "a\n\nb\n\nlate");
}
#[tokio::test]
async fn none_fragments_are_skipped() {
let composer = PromptComposer::new(
vec![
layer("a", 10, Some("alpha")),
layer("empty", 20, None),
Arc::new(FrameworkBaseLayer),
],
Vec::new(),
);
let prompt = composer.render_system_prompt(&ctx(), None).await.unwrap();
assert_eq!(prompt, "alpha");
}
#[tokio::test]
async fn framework_base_layer_is_empty() {
let composer = PromptComposer::new(vec![Arc::new(FrameworkBaseLayer)], Vec::new());
let prompt = composer.render_system_prompt(&ctx(), None).await.unwrap();
assert_eq!(prompt, "");
}
#[tokio::test]
async fn render_is_deterministic() {
let composer = PromptComposer::new(
vec![layer("a", 10, Some("alpha")), layer("b", 5, Some("beta"))],
vec![Arc::new(FixedProvider {
items: vec![item("src", "3", "text")],
})],
);
let first = composer.render_system_prompt(&ctx(), None).await.unwrap();
let second = composer.render_system_prompt(&ctx(), None).await.unwrap();
assert_eq!(first, second);
assert_eq!(first, "beta\n\nalpha\n\n## src (v3)\ntext");
}
#[tokio::test]
async fn allowed_providers_filter_drops_foreign_items() {
let composer = PromptComposer::new(
Vec::new(),
vec![Arc::new(FixedProvider {
items: vec![
ContextItem {
allowed_providers: Some(vec![ProviderId::from("anthropic")]),
..item("secret", "1", "only for anthropic")
},
item("public", "1", "for everyone"),
],
})],
);
let prompt = composer
.render_system_prompt(&ctx(), Some(&ProviderId::from("openai")))
.await
.unwrap();
assert!(!prompt.contains("only for anthropic"));
assert!(prompt.contains("## public (v1)\nfor everyone"));
}
}