166 lines
5.5 KiB
Rust
166 lines
5.5 KiB
Rust
use crate::domain::{Chunk, Role};
|
|
use crate::query::Query;
|
|
use anyhow::{anyhow, Result};
|
|
|
|
const SYSTEM_PROMPT: &str = include_str!("../../../templates/gru-mem.txt");
|
|
const BUDGET_TOTAL: usize = 32768;
|
|
const BUDGET_RESPONSE: usize = 2048;
|
|
const BUDGET_SYSTEM: usize = 400;
|
|
const BUDGET_QUESTION: usize = 150;
|
|
const BUDGET_MEMORY_MAX: usize = 1024;
|
|
const BUDGET_CHUNK_MAX: usize = 5000;
|
|
|
|
/// Builds a GRU-Mem prompt for the update gate.
|
|
pub struct PromptBuilder;
|
|
|
|
impl PromptBuilder {
|
|
/// Assemble system and user prompts for a single gate turn.
|
|
///
|
|
/// # Arguments
|
|
/// * `query` - Standing question providing the problem statement
|
|
/// * `previous_memory` - Prior memory from turn t-1, or None for t=1
|
|
/// * `chunk` - The evidence chunk to evaluate
|
|
///
|
|
/// # Returns
|
|
/// `(system_prompt, user_message)` tuple
|
|
pub fn build(query: &Query, previous_memory: Option<&str>, chunk: &Chunk) -> Result<(String, String)> {
|
|
// Render chunk as "[role] text" lines separated by blank lines
|
|
let chunk_text = Self::render_chunk(chunk)?;
|
|
let chunk_bytes = chunk_text.len();
|
|
|
|
// Memory: "No previous memory" at t=1, otherwise the given memory
|
|
let memory_text = previous_memory.unwrap_or("No previous memory");
|
|
|
|
// Check memory budget
|
|
if memory_text.len() > BUDGET_MEMORY_MAX {
|
|
return Err(anyhow!(
|
|
"Memory budget exceeded: {} > {} tokens",
|
|
memory_text.len() / 4, // rough estimate
|
|
BUDGET_MEMORY_MAX / 4
|
|
));
|
|
}
|
|
|
|
// Check chunk budget
|
|
if chunk_bytes > BUDGET_CHUNK_MAX {
|
|
return Err(anyhow!(
|
|
"Chunk budget exceeded: {} > {} bytes",
|
|
chunk_bytes,
|
|
BUDGET_CHUNK_MAX
|
|
));
|
|
}
|
|
|
|
// Assemble the user message by substituting into the template
|
|
let user_message = SYSTEM_PROMPT
|
|
.replace("{prompt}", &query.question)
|
|
.replace("{memory}", memory_text)
|
|
.replace("{chunk}", &chunk_text);
|
|
|
|
// Check total budget (rough: 4 chars ≈ 1 token)
|
|
let total_tokens = (SYSTEM_PROMPT.len() + query.question.len() + memory_text.len() + chunk_bytes) / 4;
|
|
if total_tokens + BUDGET_RESPONSE > BUDGET_TOTAL {
|
|
return Err(anyhow!(
|
|
"Total prompt budget exceeded: {} + {} (response) > {} tokens",
|
|
total_tokens,
|
|
BUDGET_RESPONSE,
|
|
BUDGET_TOTAL
|
|
));
|
|
}
|
|
|
|
Ok((SYSTEM_PROMPT.to_string(), user_message))
|
|
}
|
|
|
|
/// Render a chunk as formatted text with role labels.
|
|
fn render_chunk(chunk: &Chunk) -> Result<String> {
|
|
let mut lines = Vec::new();
|
|
|
|
for record in &chunk.records {
|
|
let role_label = match record.role {
|
|
Role::User => "[User]",
|
|
Role::Assistant => "[Assistant]",
|
|
Role::ToolResult => "[ToolResult]",
|
|
Role::System => "[System]",
|
|
};
|
|
|
|
let text = format!("{} {}", role_label, record.text);
|
|
lines.push(text);
|
|
}
|
|
|
|
Ok(lines.join("\n\n"))
|
|
}
|
|
}
|
|
|
|
#[cfg(test)]
|
|
mod tests {
|
|
use super::*;
|
|
use crate::domain::{Chunk, Record, Role, Provenance};
|
|
use time::OffsetDateTime;
|
|
|
|
#[test]
|
|
fn test_render_chunk_single_record() {
|
|
let chunk = Chunk::new(
|
|
1,
|
|
vec![
|
|
Record {
|
|
role: Role::User,
|
|
text: "Hello".to_string(),
|
|
timestamp: OffsetDateTime::now_utc(),
|
|
provenance: Provenance {
|
|
source_id: "test".to_string(),
|
|
offset: 0,
|
|
},
|
|
},
|
|
],
|
|
10,
|
|
);
|
|
|
|
let rendered = PromptBuilder::render_chunk(&chunk).unwrap();
|
|
assert!(rendered.contains("[User]"));
|
|
assert!(rendered.contains("Hello"));
|
|
}
|
|
|
|
#[test]
|
|
fn test_render_chunk_multiple_roles() {
|
|
let chunk = Chunk::new(
|
|
1,
|
|
vec![
|
|
Record {
|
|
role: Role::User,
|
|
text: "What is 2+2?".to_string(),
|
|
timestamp: OffsetDateTime::now_utc(),
|
|
provenance: Provenance {
|
|
source_id: "test".to_string(),
|
|
offset: 0,
|
|
},
|
|
},
|
|
Record {
|
|
role: Role::Assistant,
|
|
text: "The answer is 4".to_string(),
|
|
timestamp: OffsetDateTime::now_utc(),
|
|
provenance: Provenance {
|
|
source_id: "test".to_string(),
|
|
offset: 0,
|
|
},
|
|
},
|
|
Record {
|
|
role: Role::ToolResult,
|
|
text: "Tool confirmed: 4".to_string(),
|
|
timestamp: OffsetDateTime::now_utc(),
|
|
provenance: Provenance {
|
|
source_id: "test".to_string(),
|
|
offset: 0,
|
|
},
|
|
},
|
|
],
|
|
30,
|
|
);
|
|
|
|
let rendered = PromptBuilder::render_chunk(&chunk).unwrap();
|
|
assert!(rendered.contains("[User]"));
|
|
assert!(rendered.contains("[Assistant]"));
|
|
assert!(rendered.contains("[ToolResult]"));
|
|
|
|
// Check that records are separated by blank lines
|
|
assert!(rendered.contains("\n\n"));
|
|
}
|
|
}
|