Skip to content
Open
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
24 changes: 24 additions & 0 deletions src/main/java/com/google/cloud/mcp/McpToolboxClient.java
Original file line number Diff line number Diff line change
Expand Up @@ -175,6 +175,22 @@ interface Builder {
*/
Builder protocolVersion(ProtocolVersion protocolVersion);

/**
* Sets the connect timeout for the underlying HttpClient.
*
* @param connectTimeout The connect timeout.
* @return The builder instance.
*/
Builder connectTimeout(java.time.Duration connectTimeout);

/**
* Sets the request timeout for every HTTP request.
*
* @param requestTimeout The request timeout.
* @return The builder instance.
*/
Builder requestTimeout(java.time.Duration requestTimeout);

/**
* Sets a custom {@link java.net.http.HttpClient} for connection management.
*
Expand All @@ -191,6 +207,14 @@ interface Builder {
*/
Builder executor(java.util.concurrent.Executor executor);

/**
* Sets a custom Logger for telemetry and logs.
*
* @param logger The custom Logger.
* @return The builder instance.
*/
Builder logger(java.util.logging.Logger logger);

/**
* Builds and returns a new {@link McpToolboxClient} instance.
*
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -38,8 +38,11 @@ public final class McpToolboxClientBuilder implements McpToolboxClient.Builder {
private final List<ToolPreProcessor> preProcessors = new ArrayList<>();
private final List<ToolPostProcessor> postProcessors = new ArrayList<>();
private ProtocolVersion protocolVersion;
private java.time.Duration connectTimeout;
private java.time.Duration requestTimeout;
private java.net.http.HttpClient httpClient;
private java.util.concurrent.Executor executor;
private java.util.logging.Logger logger;

/** Constructs a new McpToolboxClientBuilder. */
public McpToolboxClientBuilder() {}
Expand Down Expand Up @@ -92,6 +95,18 @@ public McpToolboxClient.Builder protocolVersion(ProtocolVersion protocolVersion)
return this;
}

@Override
public McpToolboxClient.Builder connectTimeout(java.time.Duration connectTimeout) {
this.connectTimeout = connectTimeout;
return this;
}

@Override
public McpToolboxClient.Builder requestTimeout(java.time.Duration requestTimeout) {
this.requestTimeout = requestTimeout;
return this;
}

@Override
public McpToolboxClient.Builder httpClient(java.net.http.HttpClient httpClient) {
this.httpClient = httpClient;
Expand All @@ -104,6 +119,12 @@ public McpToolboxClient.Builder executor(java.util.concurrent.Executor executor)
return this;
}

@Override
public McpToolboxClient.Builder logger(java.util.logging.Logger logger) {
this.logger = logger;
return this;
}

@Override
public McpToolboxClient build() {
if (baseUrl == null || baseUrl.isEmpty()) {
Expand Down Expand Up @@ -137,7 +158,10 @@ public McpToolboxClient build() {
resolvedProvider,
this.protocolVersion,
this.httpClient,
this.executor);
this.executor,
this.connectTimeout,
this.requestTimeout,
this.logger);
return new McpToolboxClientImpl(
transport, this.headers, resolvedProvider, preProcessors, postProcessors);
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -36,29 +36,102 @@
import java.util.concurrent.CompletableFuture;
import java.util.logging.Logger;

/**
* Base class for HTTP-based MCP transports providing common functionality like session tracking,
* header merging, and credentials resolution.
*/
public abstract class BaseMcpTransport implements Transport {

/** Default static logger for BaseMcpTransport. */
protected static final Logger logger = Logger.getLogger(BaseMcpTransport.class.getName());

/** Warning message displayed when using unencrypted HTTP connections. */
protected static final String HTTP_WARNING =
"This connection is using HTTP. To prevent credential exposure, please ensure all"
+ " communication is sent over HTTPS.";

/** The base URL of the MCP service. */
protected final String baseUrl;

/** Client headers configured for the transport. */
protected final Map<String, String> clientHeaders;

/** The credentials provider for dynamic authorization. */
protected final CredentialsProvider credentialsProvider;

/** The HTTP client used for requests. */
protected final HttpClient httpClient;

/** The ObjectMapper for JSON serialization. */
protected final ObjectMapper objectMapper;

/** The preferred protocol version. */
protected final ProtocolVersion preferredProtocolVersion;

/** Lock object to synchronize initialization. */
protected final Object initLock = new Object();

/** Future indicating the status of initialization. */
protected CompletableFuture<Void> initFuture;

/** The request timeout for HTTP requests. */
protected final java.time.Duration requestTimeout;

/** The logger used for active transport logging. */
protected final Logger activeLogger;

/**
* Constructs a new BaseMcpTransport.
*
* @param baseUrl The base URL.
* @param clientHeaders The client headers.
* @param credentialsProvider The credentials provider.
* @param preferredProtocolVersion The preferred protocol version.
* @param httpClient The HTTP client.
* @param executor The executor.
*/
protected BaseMcpTransport(
final String baseUrl,
final Map<String, String> clientHeaders,
final CredentialsProvider credentialsProvider,
final ProtocolVersion preferredProtocolVersion,
final HttpClient httpClient,
final java.util.concurrent.Executor executor) {
this(
baseUrl,
clientHeaders,
credentialsProvider,
preferredProtocolVersion,
httpClient,
executor,
null,
null,
null);
}

/**
* Constructs a new BaseMcpTransport with timeouts and custom logger.
*
* @param baseUrl The base URL.
* @param clientHeaders The client headers.
* @param credentialsProvider The credentials provider.
* @param preferredProtocolVersion The preferred protocol version.
* @param httpClient The HTTP client.
* @param executor The executor.
* @param connectTimeout The connection timeout.
* @param requestTimeout The request timeout.
* @param logger The custom logger.
*/
protected BaseMcpTransport(
final String baseUrl,
final Map<String, String> clientHeaders,
final CredentialsProvider credentialsProvider,
final ProtocolVersion preferredProtocolVersion,
final HttpClient httpClient,
final java.util.concurrent.Executor executor,
final java.time.Duration connectTimeout,
final java.time.Duration requestTimeout,
final Logger logger) {
if (baseUrl == null || baseUrl.isEmpty()) {
throw new IllegalArgumentException("Base URL must be provided");
}
Expand All @@ -78,12 +151,14 @@ protected BaseMcpTransport(
HttpClient.Builder builder =
HttpClient.newBuilder()
.cookieHandler(new java.net.CookieManager())
.connectTimeout(Duration.ofSeconds(10));
.connectTimeout(connectTimeout != null ? connectTimeout : Duration.ofSeconds(10));
if (executor != null) {
builder.executor(executor);
}
this.httpClient = builder.build();
}
this.requestTimeout = requestTimeout;
this.activeLogger = logger != null ? logger : BaseMcpTransport.logger;
this.objectMapper = new ObjectMapper();
}

Expand Down Expand Up @@ -186,17 +261,29 @@ final CompletableFuture<Void> ensureInitialized(final Map<String, String> extraM
}
}

/**
* Performs the version-specific initialization handshake.
*
* @param authHeader The authorization header value, if present.
* @param handshakeHeaders The resolved headers for the handshake.
* @return A CompletableFuture that completes when initialization is done.
*/
protected abstract CompletableFuture<Void> performInitialization(
final String authHeader, final Map<String, String> handshakeHeaders);

/**
* Applies protocol-specific headers to the request builder.
*
* @param builder The HTTP request builder.
*/
protected abstract void applyProtocolHeaders(final HttpRequest.Builder builder);

@Override
public final CompletableFuture<TransportManifest> listTools(
final String toolsetName, final Map<String, String> metadata) {
if (this.baseUrl.toLowerCase(java.util.Locale.ROOT).startsWith("http://")
&& !metadata.isEmpty()) {
logger.warning(HTTP_WARNING);
activeLogger.warning(HTTP_WARNING);
}
return ensureInitialized(metadata)
.thenCompose(v -> mergeHeaders(metadata))
Expand All @@ -211,6 +298,9 @@ public final CompletableFuture<TransportManifest> listTools(
HttpRequest.newBuilder()
.uri(URI.create(url))
.POST(HttpRequest.BodyPublishers.ofString(body));
if (requestTimeout != null) {
req.timeout(requestTimeout);
}
mergedHeaders.forEach(req::setHeader);
applyProtocolHeaders(req);

Expand All @@ -230,7 +320,7 @@ public final CompletableFuture<TransportResponse> invokeTool(
final Map<String, String> metadata) {
if (this.baseUrl.toLowerCase(java.util.Locale.ROOT).startsWith("http://")
&& !metadata.isEmpty()) {
logger.warning(HTTP_WARNING);
activeLogger.warning(HTTP_WARNING);
}
return ensureInitialized(metadata)
.thenCompose(v -> mergeHeaders(metadata))
Expand All @@ -247,6 +337,9 @@ public final CompletableFuture<TransportResponse> invokeTool(
.uri(URI.create(baseUrl))
.POST(HttpRequest.BodyPublishers.ofString(requestBody));

if (requestTimeout != null) {
requestBuilder.timeout(requestTimeout);
}
mergedHeaders.forEach(requestBuilder::setHeader);
applyProtocolHeaders(requestBuilder);

Expand Down
Loading
Loading