Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
87 changes: 44 additions & 43 deletions src/args.rs
Original file line number Diff line number Diff line change
@@ -1,24 +1,15 @@
use crate::tools::{DEFAULT_MAX_BYTES, DEFAULT_MAX_LINES, DEFAULT_MAX_RESULTS, OutputMode};
use rmcp::schemars::{self, JsonSchema};
use serde::{Deserialize, Serialize};

// ---------------------------------------------------------------------------
// Default value functions for serde(default = "…")
// ---------------------------------------------------------------------------

const fn default_zero() -> usize {
0
}
const fn default_max_results() -> usize {
100
}
const fn default_max_bytes() -> usize {
5 * 1024 * 1024 // 5 MiB
fn default_max_results() -> usize {
DEFAULT_MAX_RESULTS
}
const fn default_max_lines() -> usize {
2000
fn default_max_bytes() -> usize {
DEFAULT_MAX_BYTES
}
const fn default_false() -> bool {
false
fn default_max_lines() -> usize {
DEFAULT_MAX_LINES
}
const fn default_true() -> bool {
true
Expand Down Expand Up @@ -51,14 +42,6 @@ impl StringOrVec {
}
}

fn default_file_extensions() -> Option<StringOrVec> {
None
}

fn default_output_mode() -> String {
"files_with_matches".to_string()
}

// ---------------------------------------------------------------------------
// Arg structs — serde fills in defaults automatically
// ---------------------------------------------------------------------------
Expand All @@ -69,58 +52,72 @@ fn default_output_mode() -> String {
pub struct GrepArgs {
#[schemars(description = "Directory to search in")]
pub directory: String,
#[schemars(description = "Regex pattern to search for (Rust regex; no lookaround/backrefs). Use (?i) for case-insensitive.")]
#[schemars(
description = "Regex pattern to search for (Rust regex; no lookaround/backrefs). Use (?i) for case-insensitive."
)]
pub pattern: String,
#[serde(default = "default_zero")]
#[serde(default)]
#[schemars(description = "Number of lines of context before each match (default 0)")]
pub before_context: usize,
#[serde(default = "default_zero")]
#[serde(default)]
#[schemars(description = "Number of lines of context after each match (default 0)")]
pub after_context: usize,
#[serde(default = "default_max_results")]
#[schemars(description = "Maximum number of results to return (default 100)")]
pub max_results: usize,
#[serde(default = "default_false")]
#[schemars(description = "Case-insensitive search (default false). Equivalent to prefixing pattern with (?i).")]
#[serde(default)]
#[schemars(
description = "Case-insensitive search (default false). Equivalent to prefixing pattern with (?i)."
)]
pub case_insensitive: bool,
#[serde(default = "default_false")]
#[serde(default)]
#[schemars(description = "Include hidden files and directories (default false)")]
pub include_hidden: bool,
#[serde(default = "default_false")]
#[serde(default)]
#[schemars(description = "Follow symbolic links (default false)")]
pub follow_symlinks: bool,
#[serde(default = "default_true")]
#[schemars(description = "Respect .gitignore files (default true)")]
pub respect_gitignore: bool,
#[serde(default = "default_file_extensions")]
#[schemars(description = "Restrict to files with these extensions. Accepts either a single string (\"sql\") or an array ([\"rs\", \"toml\"]). Empty means all files.")]
#[serde(default)]
#[schemars(
description = "Restrict to files with these extensions. Accepts either a single string (\"sql\") or an array ([\"rs\", \"toml\"]). Empty means all files."
)]
pub file_extensions: Option<StringOrVec>,
#[serde(default = "default_max_bytes")]
#[schemars(description = "Hard cap on total response size in bytes (default ~5 MiB). Truncates with a marker.")]
#[schemars(
description = "Hard cap on total response size in bytes (default ~5 MiB). Truncates with a marker."
)]
pub max_bytes: usize,
#[serde(default = "default_output_mode")]
#[schemars(description = "Output mode: 'files_with_matches' (default — list file paths only), 'content' (matching lines with line numbers), 'count' (per-file match tallies as path: N).")]
pub output_mode: String,
#[serde(default)]
#[schemars(
description = "Output mode: 'files_with_matches' (default — list file paths only), 'content' (matching lines with line numbers), 'count' (per-file match tallies as path: N)."
)]
pub output_mode: OutputMode,
}

#[derive(Debug, Serialize, Deserialize, JsonSchema)]
#[serde(rename_all = "snake_case")]
pub struct FindArgs {
#[schemars(description = "Directory to search in")]
pub directory: String,
#[schemars(description = "Regex pattern to match against filenames (Rust regex; no lookaround/backrefs)")]
#[schemars(
description = "Regex pattern to match against filenames (Rust regex; no lookaround/backrefs)"
)]
pub pattern: String,
#[serde(default = "default_max_results")]
#[schemars(description = "Maximum number of results to return (default 100)")]
pub max_results: usize,
#[serde(default = "default_false")]
#[serde(default)]
#[schemars(description = "Include hidden files and directories (default false)")]
pub include_hidden: bool,
#[serde(default = "default_true")]
#[schemars(description = "Respect .gitignore files (default true)")]
pub respect_gitignore: bool,
#[serde(default = "default_true")]
#[schemars(description = "Match the basename only (default true). Set false to match the full path.")]
#[schemars(
description = "Match the basename only (default true). Set false to match the full path."
)]
pub match_basename: bool,
}

Expand All @@ -129,8 +126,10 @@ pub struct FindArgs {
pub struct CatArgs {
#[schemars(description = "Path to the file to read")]
pub file_path: String,
#[serde(default = "default_zero")]
#[schemars(description = "Line offset to start from (0-based, default 0). Use to paginate long files.")]
#[serde(default)]
#[schemars(
description = "Line offset to start from (0-based, default 0). Use to paginate long files."
)]
pub offset: usize,
#[serde(default = "default_max_lines")]
#[schemars(description = "Maximum number of lines to return (default 2000)")]
Expand All @@ -144,6 +143,8 @@ pub struct CatArgs {
#[serde(rename_all = "snake_case")]
pub struct MemoriesArgs {
#[serde(default)]
#[schemars(description = "Optional memory file name (relative to memory dir, e.g. \"user_role.md\"). If omitted, returns the index from MEMORY.md or a directory listing.")]
#[schemars(
description = "Optional memory file name (relative to memory dir, e.g. \"user_role.md\"). If omitted, returns the index from MEMORY.md or a directory listing."
)]
pub name: Option<String>,
}
7 changes: 5 additions & 2 deletions src/cli.rs
Original file line number Diff line number Diff line change
@@ -1,14 +1,17 @@
use std::path::PathBuf;
use clap::Parser;
use std::net::SocketAddr;
use std::path::PathBuf;

/// Command-line arguments for `code-mcp`.
///
/// Parsed via clap. All fields are `pub(crate)` because they're only consumed
/// by `main`; the struct itself is `pub` so it can be referenced from other
/// crate modules.
#[derive(Debug, Parser)]
#[command(name = "code-mcp", about = "Streamable HTTP MCP server for code search/read tools")]
#[command(
name = "code-mcp",
about = "Streamable HTTP MCP server for code search/read tools"
)]
pub struct Args {
/// Address to bind, e.g. 0.0.0.0:8080
#[arg(long, default_value = "0.0.0.0:8080")]
Expand Down
12 changes: 4 additions & 8 deletions src/error.rs
Original file line number Diff line number Diff line change
Expand Up @@ -71,15 +71,11 @@ impl From<AppError> for ErrorData {
| AppError::GrepRegex(_)
| AppError::InvalidRequest(_)
| AppError::NotFound(_)
| AppError::OutOfScope(_) => ErrorData::invalid_params(
"invalid_params",
Some(json!({"error": err.to_string()})),
),
| AppError::OutOfScope(_) => {
ErrorData::invalid_params("invalid_params", Some(json!({"error": err.to_string()})))
}
AppError::Io(_) | AppError::Ignore(_) | AppError::Internal(_) | AppError::Axum(_) => {
ErrorData::internal_error(
"internal_error",
Some(json!({"error": err.to_string()})),
)
ErrorData::internal_error("internal_error", Some(json!({"error": err.to_string()})))
}
}
}
Expand Down
58 changes: 44 additions & 14 deletions src/gate.rs
Original file line number Diff line number Diff line change
Expand Up @@ -6,9 +6,7 @@ use axum::extract::{ConnectInfo, State};
use axum::http::{HeaderMap, Method, Request, StatusCode, header};
use axum::middleware::Next;
use axum::response::{IntoResponse, Response};
use rmcp::transport::streamable_http_server::session::{
SessionId, local::LocalSessionManager,
};
use rmcp::transport::streamable_http_server::session::{SessionId, local::LocalSessionManager};

use crate::limiter::PeerLimiter;
use crate::reaper::ActivityTracker;
Expand Down Expand Up @@ -107,7 +105,7 @@ fn peer_ip(headers: &HeaderMap, addr: SocketAddr, trust_xff: bool) -> IpAddr {
.map(str::trim)
.and_then(|s| s.parse::<IpAddr>().ok())
{
tracing::info!(peer = %ip, socket = %addr.ip(), "resolved peer IP from X-Forwarded-For");
tracing::debug!(peer = %ip, socket = %addr.ip(), "resolved peer IP from X-Forwarded-For");
return ip;
}
addr.ip()
Expand Down Expand Up @@ -155,7 +153,10 @@ mod tests {
activity: Arc::new(ActivityTracker::new()),
});
let app = build_app(ctx);
let res = app.oneshot(req(Method::GET, dummy_addr(), false)).await.unwrap();
let res = app
.oneshot(req(Method::GET, dummy_addr(), false))
.await
.unwrap();
assert_eq!(res.status(), StatusCode::OK);
}

Expand All @@ -170,7 +171,10 @@ mod tests {
activity: Arc::new(ActivityTracker::new()),
});
let app = build_app(ctx);
let res = app.oneshot(req(Method::POST, dummy_addr(), true)).await.unwrap();
let res = app
.oneshot(req(Method::POST, dummy_addr(), true))
.await
.unwrap();
assert_eq!(res.status(), StatusCode::OK);
}

Expand All @@ -184,7 +188,10 @@ mod tests {
activity: Arc::new(ActivityTracker::new()),
});
let app = build_app(ctx);
let res = app.oneshot(req(Method::POST, dummy_addr(), false)).await.unwrap();
let res = app
.oneshot(req(Method::POST, dummy_addr(), false))
.await
.unwrap();
assert_eq!(res.status(), StatusCode::SERVICE_UNAVAILABLE);
assert_eq!(res.headers().get(header::RETRY_AFTER).unwrap(), "5");
}
Expand All @@ -201,31 +208,48 @@ mod tests {
let app = build_app(ctx);

for _ in 0..2 {
let res = app.clone().oneshot(req(Method::POST, dummy_addr(), false)).await.unwrap();
let res = app
.clone()
.oneshot(req(Method::POST, dummy_addr(), false))
.await
.unwrap();
assert_eq!(res.status(), StatusCode::OK);
}
let res = app.clone().oneshot(req(Method::POST, dummy_addr(), false)).await.unwrap();
let res = app
.clone()
.oneshot(req(Method::POST, dummy_addr(), false))
.await
.unwrap();
assert_eq!(res.status(), StatusCode::TOO_MANY_REQUESTS);
assert!(res.headers().get(header::RETRY_AFTER).is_some());

// A different peer is unaffected.
let res = app.oneshot(req(Method::POST, other_addr(), false)).await.unwrap();
let res = app
.oneshot(req(Method::POST, other_addr(), false))
.await
.unwrap();
assert_eq!(res.status(), StatusCode::OK);
}

#[test]
fn peer_ip_uses_socket_addr_by_default() {
let h = HeaderMap::new();
let addr: SocketAddr = "10.0.0.5:1234".parse().unwrap();
assert_eq!(peer_ip(&h, addr, false), IpAddr::V4(Ipv4Addr::new(10, 0, 0, 5)));
assert_eq!(
peer_ip(&h, addr, false),
IpAddr::V4(Ipv4Addr::new(10, 0, 0, 5))
);
}

#[test]
fn peer_ip_ignores_xff_when_untrusted() {
let mut h = HeaderMap::new();
h.insert("x-forwarded-for", "1.2.3.4".parse().unwrap());
let addr: SocketAddr = "10.0.0.5:1234".parse().unwrap();
assert_eq!(peer_ip(&h, addr, false), IpAddr::V4(Ipv4Addr::new(10, 0, 0, 5)));
assert_eq!(
peer_ip(&h, addr, false),
IpAddr::V4(Ipv4Addr::new(10, 0, 0, 5))
);
}

#[test]
Expand All @@ -236,14 +260,20 @@ mod tests {
let mut h = HeaderMap::new();
h.insert("x-forwarded-for", "1.2.3.4, 5.6.7.8".parse().unwrap());
let addr: SocketAddr = "10.0.0.5:1234".parse().unwrap();
assert_eq!(peer_ip(&h, addr, true), IpAddr::V4(Ipv4Addr::new(5, 6, 7, 8)));
assert_eq!(
peer_ip(&h, addr, true),
IpAddr::V4(Ipv4Addr::new(5, 6, 7, 8))
);
}

#[test]
fn peer_ip_falls_back_to_socket_when_xff_unparseable() {
let mut h = HeaderMap::new();
h.insert("x-forwarded-for", "garbage".parse().unwrap());
let addr: SocketAddr = "10.0.0.5:1234".parse().unwrap();
assert_eq!(peer_ip(&h, addr, true), IpAddr::V4(Ipv4Addr::new(10, 0, 0, 5)));
assert_eq!(
peer_ip(&h, addr, true),
IpAddr::V4(Ipv4Addr::new(10, 0, 0, 5))
);
}
}
4 changes: 3 additions & 1 deletion src/limiter.rs
Original file line number Diff line number Diff line change
Expand Up @@ -102,7 +102,9 @@ mod tests {
let l = PeerLimiter::new(2.0, 0.001, 1024);
l.try_consume(ip(1)).unwrap();
l.try_consume(ip(1)).unwrap();
let err = l.try_consume(ip(1)).expect_err("third call should be rate-limited");
let err = l
.try_consume(ip(1))
.expect_err("third call should be rate-limited");
assert!(err > Duration::from_secs(0));
}

Expand Down
10 changes: 4 additions & 6 deletions src/main.rs
Original file line number Diff line number Diff line change
Expand Up @@ -39,6 +39,7 @@
//! `invalid_params`. See <https://modelcontextprotocol.io/docs/tutorials/security/authorization>.

mod args;
mod cli;
mod error;
mod gate;
mod limiter;
Expand All @@ -47,26 +48,23 @@ mod reaper;
mod scope;
mod server;
mod tools;
mod cli;

use clap::Parser;
use std::sync::Arc;
use std::time::Duration;
use clap::Parser;

use std::net::SocketAddr;


use rmcp::transport::streamable_http_server::{
StreamableHttpServerConfig, StreamableHttpService,
session::local::LocalSessionManager,
StreamableHttpServerConfig, StreamableHttpService, session::local::LocalSessionManager,
};
use tokio_util::sync::CancellationToken;

use crate::cli::Args;
use crate::gate::{GateCtx, gate};
use crate::limiter::PeerLimiter;
use crate::scope::Scope;
use crate::server::CodeMcpServer;
use crate::cli::Args;

use crate::error::AppError;

Expand Down
Loading