diff --git a/crates/goose/src/config/signup_openrouter/server.rs b/crates/goose/src/config/signup_openrouter/server.rs
index d809382f6..027fe70fa 100644
--- a/crates/goose/src/config/signup_openrouter/server.rs
+++ b/crates/goose/src/config/signup_openrouter/server.rs
@@ -54,8 +54,8 @@ async fn handle_callback(
.contents_utf8()
.expect("error.html is not valid UTF-8");
- env.add_template("error", template_content).unwrap();
- let tmpl = env.get_template("error").unwrap();
+ env.add_template("error.html", template_content).unwrap();
+ let tmpl = env.get_template("error.html").unwrap();
let rendered = tmpl.render(context! { error => error }).unwrap();
return (StatusCode::BAD_REQUEST, Html(rendered));
@@ -84,3 +84,43 @@ async fn handle_callback(
(StatusCode::BAD_REQUEST, Html(invalid_html.to_string()))
}
+
+#[cfg(test)]
+mod tests {
+ use super::*;
+ use axum::body::to_bytes;
+
+ async fn error_response(error: &str) -> (StatusCode, String) {
+ let state = std::sync::Arc::new(tokio::sync::Mutex::new(None));
+ let response = handle_callback(
+ Query(CallbackQuery {
+ code: None,
+ error: Some(error.to_string()),
+ }),
+ axum::extract::State(state),
+ )
+ .await
+ .into_response();
+ let status = response.status();
+ let body = to_bytes(response.into_body(), usize::MAX).await.unwrap();
+ (status, String::from_utf8(body.to_vec()).unwrap())
+ }
+
+ #[tokio::test]
+ async fn error_response_escapes_html() {
+ let payload = r#"&"#;
+ let (status, body) = error_response(payload).await;
+
+ assert_eq!(status, StatusCode::BAD_REQUEST);
+ assert!(!body.contains(payload));
+ assert!(body.contains("<script>"));
+ assert!(body.contains("&"));
+ }
+
+ #[tokio::test]
+ async fn error_response_preserves_plain_text() {
+ let (_, body) = error_response("authorization denied").await;
+
+ assert!(body.contains("authorization denied"));
+ }
+}
diff --git a/crates/goose/src/config/signup_tetrate/server.rs b/crates/goose/src/config/signup_tetrate/server.rs
index bf72c92e2..82b54c157 100644
--- a/crates/goose/src/config/signup_tetrate/server.rs
+++ b/crates/goose/src/config/signup_tetrate/server.rs
@@ -51,8 +51,8 @@ async fn handle_callback(
.contents_utf8()
.expect("error.html is not valid UTF-8");
- env.add_template("error", template_content).unwrap();
- let tmpl = env.get_template("error").unwrap();
+ env.add_template("error.html", template_content).unwrap();
+ let tmpl = env.get_template("error.html").unwrap();
let rendered = tmpl.render(context! { error => error }).unwrap();
return (StatusCode::BAD_REQUEST, Html(rendered));
@@ -81,3 +81,43 @@ async fn handle_callback(
(StatusCode::BAD_REQUEST, Html(invalid_html.to_string()))
}
+
+#[cfg(test)]
+mod tests {
+ use super::*;
+ use axum::body::to_bytes;
+
+ async fn error_response(error: &str) -> (StatusCode, String) {
+ let state = std::sync::Arc::new(tokio::sync::Mutex::new(None));
+ let response = handle_callback(
+ Query(CallbackQuery {
+ code: None,
+ error: Some(error.to_string()),
+ }),
+ axum::extract::State(state),
+ )
+ .await
+ .into_response();
+ let status = response.status();
+ let body = to_bytes(response.into_body(), usize::MAX).await.unwrap();
+ (status, String::from_utf8(body.to_vec()).unwrap())
+ }
+
+ #[tokio::test]
+ async fn error_response_escapes_html() {
+ let payload = r#"&"#;
+ let (status, body) = error_response(payload).await;
+
+ assert_eq!(status, StatusCode::BAD_REQUEST);
+ assert!(!body.contains(payload));
+ assert!(body.contains("<script>"));
+ assert!(body.contains("&"));
+ }
+
+ #[tokio::test]
+ async fn error_response_preserves_plain_text() {
+ let (_, body) = error_response("authorization denied").await;
+
+ assert!(body.contains("authorization denied"));
+ }
+}
diff --git a/crates/goose/src/oauth/mod.rs b/crates/goose/src/oauth/mod.rs
index 09a52066e..07b7214fa 100644
--- a/crates/goose/src/oauth/mod.rs
+++ b/crates/goose/src/oauth/mod.rs
@@ -6,7 +6,7 @@ use axum::extract::{Query, State};
use axum::response::Html;
use axum::routing::get;
use axum::Router;
-use minijinja::render;
+use minijinja::{context, Environment};
use oauth2::{Scope, TokenResponse};
use rmcp::transport::auth::{
AuthError, AuthorizationRequest, CredentialStore, OAuthClientConfig, OAuthState,
@@ -62,6 +62,16 @@ fn oauth_callback_timeout() -> Duration {
resolve_oauth_callback_timeout(timeout.as_deref())
}
+fn render_oauth_callback(name: &str) -> String {
+ Environment::new()
+ .render_named_str(
+ "oauth_callback.html",
+ CALLBACK_TEMPLATE,
+ context! { name => name },
+ )
+ .expect("failed to render OAuth callback")
+}
+
fn announce_authorization_url(name: &str, authorization_url: &str) {
warn!(
"[OAuth:{}] If the browser did not open, authorize manually at: {}",
@@ -380,7 +390,7 @@ pub async fn oauth_flow_with_challenge(
let app_state = AppState {
callback_receiver: Arc::new(Mutex::new(Some(callback_sender))),
};
- let rendered = render!(CALLBACK_TEMPLATE, name => name);
+ let rendered = render_oauth_callback(name);
let handler = move |Query(params): Query, State(state): State| {
let rendered = rendered.clone();
async move {
@@ -512,6 +522,24 @@ mod tests {
);
}
+ #[test]
+ fn oauth_callback_escapes_extension_name() {
+ let payload = r#"&"#;
+ let rendered = render_oauth_callback(payload);
+
+ assert!(!rendered.contains(payload));
+ assert!(rendered.contains("<script>"));
+ assert!(rendered.contains("&"));
+ }
+
+ #[test]
+ fn oauth_callback_preserves_plain_extension_name() {
+ let rendered = render_oauth_callback("Example MCP");
+
+ assert!(rendered.contains("Example MCP OAuth Success"));
+ assert!(rendered.contains(">Example MCP"));
+ }
+
#[tokio::test]
async fn wait_for_callback_returns_received_callback_url() {
let (sender, receiver) = oneshot::channel();