Cors and token (#5850)
Co-authored-by: Douwe Osinga <douwe@squareup.com>
This commit is contained in:
@@ -3,7 +3,7 @@ use axum::response::Redirect;
|
||||
use axum::{
|
||||
extract::{
|
||||
ws::{Message, WebSocket, WebSocketUpgrade},
|
||||
Request, State,
|
||||
Query, Request, State,
|
||||
},
|
||||
http::StatusCode,
|
||||
middleware::{self, Next},
|
||||
@@ -21,7 +21,7 @@ use serde::{Deserialize, Serialize};
|
||||
use serde_json::Value;
|
||||
use std::{net::SocketAddr, sync::Arc};
|
||||
use tokio::sync::{Mutex, RwLock};
|
||||
use tower_http::cors::{Any, CorsLayer};
|
||||
use tower_http::cors::{AllowOrigin, Any, CorsLayer};
|
||||
use tracing::error;
|
||||
use webbrowser;
|
||||
|
||||
@@ -32,6 +32,7 @@ struct AppState {
|
||||
agent: Arc<Agent>,
|
||||
cancellations: CancellationStore,
|
||||
auth_token: Option<String>,
|
||||
ws_token: String,
|
||||
}
|
||||
|
||||
#[derive(Serialize, Deserialize)]
|
||||
@@ -87,17 +88,14 @@ async fn auth_middleware(
|
||||
req: Request,
|
||||
next: Next,
|
||||
) -> Result<Response, StatusCode> {
|
||||
// Skip auth for health check
|
||||
if req.uri().path() == "/api/health" {
|
||||
return Ok(next.run(req).await);
|
||||
}
|
||||
|
||||
// If no auth token is configured, skip authentication entirely
|
||||
let Some(ref expected_token) = state.auth_token else {
|
||||
return Ok(next.run(req).await);
|
||||
};
|
||||
|
||||
// Check for Bearer token first
|
||||
if let Some(auth_header) = req.headers().get("authorization") {
|
||||
if let Ok(auth_str) = auth_header.to_str() {
|
||||
if let Some(token) = auth_str.strip_prefix("Bearer ") {
|
||||
@@ -106,7 +104,6 @@ async fn auth_middleware(
|
||||
}
|
||||
}
|
||||
|
||||
// Check for Basic auth (password-only, ignore username)
|
||||
if let Some(basic_token) = auth_str.strip_prefix("Basic ") {
|
||||
if let Ok(decoded) = base64::engine::general_purpose::STANDARD.decode(basic_token) {
|
||||
if let Ok(credentials) = String::from_utf8(decoded) {
|
||||
@@ -119,7 +116,6 @@ async fn auth_middleware(
|
||||
}
|
||||
}
|
||||
|
||||
// Authentication failed - return 401 with WWW-Authenticate header
|
||||
let mut response = Response::new("Authentication required".into());
|
||||
*response.status_mut() = StatusCode::UNAUTHORIZED;
|
||||
response.headers_mut().insert(
|
||||
@@ -135,7 +131,6 @@ pub async fn handle_web(
|
||||
open: bool,
|
||||
auth_token: Option<String>,
|
||||
) -> Result<()> {
|
||||
// Setup logging
|
||||
crate::logging::setup_logging(Some("goose-web"), None)?;
|
||||
|
||||
let config = goose::config::Config::global();
|
||||
@@ -176,10 +171,34 @@ pub async fn handle_web(
|
||||
}
|
||||
}
|
||||
|
||||
let ws_token = if auth_token.is_none() {
|
||||
uuid::Uuid::new_v4().to_string()
|
||||
} else {
|
||||
String::new()
|
||||
};
|
||||
|
||||
let state = AppState {
|
||||
agent: Arc::new(agent),
|
||||
cancellations: Arc::new(RwLock::new(std::collections::HashMap::new())),
|
||||
auth_token,
|
||||
auth_token: auth_token.clone(),
|
||||
ws_token,
|
||||
};
|
||||
|
||||
let cors_layer = if auth_token.is_none() {
|
||||
let allowed_origins = [
|
||||
"http://localhost:3000".parse().unwrap(),
|
||||
"http://127.0.0.1:3000".parse().unwrap(),
|
||||
format!("http://{}:{}", host, port).parse().unwrap(),
|
||||
];
|
||||
CorsLayer::new()
|
||||
.allow_origin(AllowOrigin::list(allowed_origins))
|
||||
.allow_methods(Any)
|
||||
.allow_headers(Any)
|
||||
} else {
|
||||
CorsLayer::new()
|
||||
.allow_origin(Any)
|
||||
.allow_methods(Any)
|
||||
.allow_headers(Any)
|
||||
};
|
||||
|
||||
let app = Router::new()
|
||||
@@ -194,12 +213,7 @@ pub async fn handle_web(
|
||||
state.clone(),
|
||||
auth_middleware,
|
||||
))
|
||||
.layer(
|
||||
CorsLayer::new()
|
||||
.allow_origin(Any)
|
||||
.allow_methods(Any)
|
||||
.allow_headers(Any),
|
||||
)
|
||||
.layer(cors_layer)
|
||||
.with_state(state);
|
||||
|
||||
let addr: SocketAddr = format!("{}:{}", host, port).parse()?;
|
||||
@@ -214,7 +228,6 @@ pub async fn handle_web(
|
||||
println!(" Press Ctrl+C to stop\n");
|
||||
|
||||
if open {
|
||||
// Open browser
|
||||
let url = format!("http://{}", addr);
|
||||
if let Err(e) = webbrowser::open(&url) {
|
||||
eprintln!("Failed to open browser: {}", e);
|
||||
@@ -241,14 +254,15 @@ async fn serve_index() -> Result<Redirect, (http::StatusCode, String)> {
|
||||
|
||||
async fn serve_session(
|
||||
axum::extract::Path(session_name): axum::extract::Path<String>,
|
||||
State(state): State<AppState>,
|
||||
) -> Html<String> {
|
||||
let html = include_str!("../../static/index.html");
|
||||
// Inject the session name into the HTML so JavaScript can use it
|
||||
let html_with_session = html.replace(
|
||||
"<script src=\"/static/script.js\"></script>",
|
||||
&format!(
|
||||
"<script>window.GOOSE_SESSION_NAME = '{}';</script>\n <script src=\"/static/script.js\"></script>",
|
||||
session_name
|
||||
"<script>window.GOOSE_SESSION_NAME = '{}'; window.GOOSE_WS_TOKEN = '{}';</script>\n <script src=\"/static/script.js\"></script>",
|
||||
session_name,
|
||||
state.ws_token
|
||||
)
|
||||
);
|
||||
Html(html_with_session)
|
||||
@@ -324,11 +338,25 @@ async fn get_session(
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct WsQuery {
|
||||
token: Option<String>,
|
||||
}
|
||||
|
||||
async fn websocket_handler(
|
||||
ws: WebSocketUpgrade,
|
||||
State(state): State<AppState>,
|
||||
) -> impl IntoResponse {
|
||||
ws.on_upgrade(|socket| handle_socket(socket, state))
|
||||
Query(query): Query<WsQuery>,
|
||||
) -> Result<impl IntoResponse, StatusCode> {
|
||||
if state.auth_token.is_none() {
|
||||
let provided_token = query.token.as_deref().unwrap_or("");
|
||||
if provided_token != state.ws_token {
|
||||
tracing::warn!("WebSocket connection rejected: invalid token");
|
||||
return Err(StatusCode::FORBIDDEN);
|
||||
}
|
||||
}
|
||||
|
||||
Ok(ws.on_upgrade(|socket| handle_socket(socket, state)))
|
||||
}
|
||||
|
||||
async fn handle_socket(socket: WebSocket, state: AppState) {
|
||||
|
||||
@@ -138,7 +138,8 @@ function removeThinkingIndicator() {
|
||||
// Connect to WebSocket
|
||||
function connectWebSocket() {
|
||||
const protocol = window.location.protocol === 'https:' ? 'wss:' : 'ws:';
|
||||
const wsUrl = `${protocol}//${window.location.host}/ws`;
|
||||
const token = window.GOOSE_WS_TOKEN || '';
|
||||
const wsUrl = `${protocol}//${window.location.host}/ws?token=${encodeURIComponent(token)}`;
|
||||
|
||||
socket = new WebSocket(wsUrl);
|
||||
|
||||
@@ -520,4 +521,4 @@ function updateSessionTitle() {
|
||||
}
|
||||
|
||||
// Update title on load
|
||||
updateSessionTitle();
|
||||
updateSessionTitle();
|
||||
|
||||
Reference in New Issue
Block a user