feat: Support extending the system prompt (#1167)

This commit is contained in:
Bradley Axen
2025-02-11 20:03:32 -08:00
committed by GitHub
parent a5e2419380
commit 6220ef054f
8 changed files with 88 additions and 1 deletions
+35
View File
@@ -17,6 +17,16 @@ struct VersionsResponse {
default_version: String,
}
#[derive(Deserialize)]
struct ExtendPromptRequest {
extension: String,
}
#[derive(Serialize)]
struct ExtendPromptResponse {
success: bool,
}
#[derive(Deserialize)]
struct CreateAgentRequest {
version: Option<String>,
@@ -61,6 +71,30 @@ async fn get_versions() -> Json<VersionsResponse> {
})
}
async fn extend_prompt(
State(state): State<AppState>,
headers: HeaderMap,
Json(payload): Json<ExtendPromptRequest>,
) -> Result<Json<ExtendPromptResponse>, StatusCode> {
// Verify secret key
let secret_key = headers
.get("X-Secret-Key")
.and_then(|value| value.to_str().ok())
.ok_or(StatusCode::UNAUTHORIZED)?;
if secret_key != state.secret_key {
return Err(StatusCode::UNAUTHORIZED);
}
let mut agent = state.agent.lock().await;
if let Some(ref mut agent) = *agent {
agent.extend_system_prompt(payload.extension).await;
Ok(Json(ExtendPromptResponse { success: true }))
} else {
Err(StatusCode::NOT_FOUND)
}
}
async fn create_agent(
State(state): State<AppState>,
headers: HeaderMap,
@@ -132,6 +166,7 @@ pub fn routes(state: AppState) -> Router {
Router::new()
.route("/agent/versions", get(get_versions))
.route("/agent/providers", get(list_providers))
.route("/agent/prompt", post(extend_prompt))
.route("/agent", post(create_agent))
.with_state(state)
}