use crate::classify::tiers::ClassificationResult;
#[allow(unused_imports)]
use crate::classify::tiers::llm::{LlmVerdict, SYSTEM_PROMPT};
pub struct BedrockClassifier {
#[allow(dead_code)] pub(crate) model: String,
#[cfg(feature = "bedrock")]
inner: trusty_common::inference::BedrockAdapter,
}
pub const DEFAULT_BEDROCK_MODEL: &str = "anthropic.claude-3-haiku-20240307-v1:0";
impl BedrockClassifier {
#[cfg(feature = "bedrock")]
pub async fn new(model: &str) -> Result<Self, String> {
Ok(Self {
model: model.to_string(),
inner: trusty_common::inference::BedrockAdapter::new(None),
})
}
#[cfg(feature = "bedrock")]
pub async fn with_region(model: &str, region: Option<&str>) -> Result<Self, String> {
Ok(Self {
model: model.to_string(),
inner: trusty_common::inference::BedrockAdapter::new(region),
})
}
#[cfg(not(feature = "bedrock"))]
pub async fn new(_model: &str) -> Result<Self, String> {
Err("bedrock feature not compiled in — rebuild with --features bedrock".to_string())
}
#[cfg(not(feature = "bedrock"))]
pub async fn with_region(_model: &str, _region: Option<&str>) -> Result<Self, String> {
Err("bedrock feature not compiled in — rebuild with --features bedrock".to_string())
}
#[cfg(feature = "bedrock")]
pub async fn classify_batch_bedrock(
&self,
messages: &[&str],
) -> Vec<Option<ClassificationResult>> {
let mut out = Vec::with_capacity(messages.len());
for msg in messages {
out.push(self.classify_one(msg).await);
}
out
}
#[cfg(not(feature = "bedrock"))]
pub async fn classify_batch_bedrock(
&self,
messages: &[&str],
) -> Vec<Option<ClassificationResult>> {
vec![None; messages.len()]
}
#[cfg(feature = "bedrock")]
async fn classify_one(&self, message: &str) -> Option<ClassificationResult> {
use crate::core::models::ClassificationMethod;
use tracing::warn;
use trusty_common::inference::{ChatMessage, ChatRequest, InferenceAdapter};
let mut req = ChatRequest::new(
self.model.clone(),
vec![
ChatMessage::system(SYSTEM_PROMPT),
ChatMessage::user(format!("Classify this commit message:\n\n{message}")),
],
);
req.temperature = Some(0.0);
req.max_tokens = Some(256);
let resp = match self.inner.chat(&req).await {
Ok(r) => r,
Err(e) => {
warn!(error = %e, "bedrock converse call failed");
return None;
}
};
let text = resp.first_text().unwrap_or_default();
let verdict: LlmVerdict = match serde_json::from_str(text.trim()) {
Ok(v) => v,
Err(e) => {
warn!(error = %e, raw = %text, "bedrock verdict parse failed");
return None;
}
};
Some(ClassificationResult {
category: verdict.category,
subcategory: verdict.subcategory,
top_level: None,
confidence: verdict.confidence.clamp(0.0, 1.0),
method: ClassificationMethod::LlmFallback,
ticket_id: None,
complexity: verdict.complexity.map(|v| v.clamp(1, 5)),
})
}
}
#[cfg(all(test, not(feature = "bedrock")))]
mod tests {
use super::*;
#[tokio::test]
async fn bedrock_stub_returns_error_without_feature() {
let result = BedrockClassifier::new("anthropic.claude-3-haiku-20240307-v1:0").await;
let err = match result {
Err(e) => e,
Ok(_) => panic!("must error without feature"),
};
assert!(err.contains("bedrock feature not compiled in"));
}
#[test]
fn shared_system_prompt_contains_complexity_instruction() {
assert!(
SYSTEM_PROMPT.contains("complexity"),
"shared SYSTEM_PROMPT must instruct the model to return a complexity score"
);
}
}