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
Original file line number Diff line number Diff line change
Expand Up @@ -93,6 +93,7 @@ enum CapabilitySchemaBuilder {
"organization": stringSchema("Optional organization name."),
"workspace": stringSchema("Optional workspace name."),
"repository": stringSchema("Optional owner/repository identifier."),
"task": stringSchema("Optional stable task identifier for task-scoped assignments."),
"current_skill_ids": InputSchema(
type: "array",
description: "Stable IDs for skills already active in the agent session.",
Expand Down Expand Up @@ -224,6 +225,7 @@ enum CapabilitySchemaBuilder {
"organization": stringSchema("Optional organization name."),
"workspace": stringSchema("Optional workspace name."),
"repository": stringSchema("Optional owner/repository identifier."),
"task": stringSchema("Optional stable task identifier for task-scoped assignments."),
],
additionalProperties: false
)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -59,10 +59,6 @@ enum SkillRuntimeToolHandlers {
guard let request = nonempty(string(arguments["request"])) else {
throw Abort(.badRequest, reason: "request is required")
}
let contextArgument = arguments["context"]
guard contextArgument == nil || decode(RuntimeContext.self, contextArgument) != nil else {
throw Abort(.badRequest, reason: "context must be an object with string identity fields")
}
let currentSkillArgument = arguments["current_skill_ids"]
guard currentSkillArgument == nil || decode([String].self, currentSkillArgument) != nil else {
throw Abort(.badRequest, reason: "current_skill_ids must be an array of strings")
Expand All @@ -72,11 +68,7 @@ enum SkillRuntimeToolHandlers {
throw Abort(.badRequest, reason: "available_tools must be an array of structured tool descriptions")
}
let event = try optionalBoundedString(arguments["event"], field: "event", max: 128)
var context = decode(RuntimeContext.self, contextArgument) ?? .init()
context.user = nonempty(string(arguments["user"])) ?? context.user
context.organization = nonempty(string(arguments["organization"])) ?? context.organization
context.workspace = nonempty(string(arguments["workspace"])) ?? context.workspace
context.repository = nonempty(string(arguments["repository"])) ?? context.repository
let context = try runtimeContext(arguments: arguments)
let response = try await SkillRuntimeResolver.resolve(
projectId: projectId,
request: request,
Expand All @@ -89,6 +81,20 @@ enum SkillRuntimeToolHandlers {
return output(response)
}

static func runtimeContext(arguments: [String: JSONValue]) throws -> RuntimeContext {
let contextArgument = arguments["context"]
guard contextArgument == nil || decode(RuntimeContext.self, contextArgument) != nil else {
throw Abort(.badRequest, reason: "context must be an object with string identity fields")
}
var context = decode(RuntimeContext.self, contextArgument) ?? .init()
context.user = nonempty(string(arguments["user"])) ?? context.user
context.organization = nonempty(string(arguments["organization"])) ?? context.organization
context.workspace = nonempty(string(arguments["workspace"])) ?? context.workspace
context.repository = nonempty(string(arguments["repository"])) ?? context.repository
context.task = nonempty(string(arguments["task"])) ?? context.task
return context
}

private static func legacyDiscover(
_ arguments: [String: JSONValue],
db: Database,
Expand Down
73 changes: 53 additions & 20 deletions services/mcp-gateway/Sources/App/Sync/RepoFetcher.swift
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,8 @@ struct RepoFetchOutcome {
}

struct RepoFetcher {
typealias TarballDataLoader = @Sendable (URLRequest, String?) async throws -> (Data, URLResponse)

let app: Application

func fetch(owner: String, repo: String, ref: String, token: String? = nil) async throws -> RepoFetchOutcome {
Expand All @@ -23,27 +25,16 @@ struct RepoFetcher {
throw RepoFetcherError.invalidURL
}

var request = URLRequest(url: url)
request.httpMethod = "GET"
request.timeoutInterval = Self.fetchTimeoutSeconds()
let authToken = token ?? Environment.get("GITHUB_TOKEN")
if let authToken = authToken {
request.setValue("Bearer \(authToken)", forHTTPHeaderField: "Authorization")
}
request.setValue("application/vnd.github.v3+json", forHTTPHeaderField: "Accept")

let session: URLSession
if let authToken = authToken {
let delegate = GitHubTarballRedirectDelegate(bearerToken: authToken)
session = URLSession(configuration: .default, delegate: delegate, delegateQueue: nil)
} else {
session = URLSession.shared
}
let (data, response) = try await session.data(for: request)

guard let httpResponse = response as? HTTPURLResponse, httpResponse.statusCode == 200 else {
let code = (response as? HTTPURLResponse)?.statusCode ?? 0
throw RepoFetcherError.fetchFailed(status: code)
let data = try await Self.downloadTarball(url: url, authToken: authToken) { request, bearerToken in
let session: URLSession
if let bearerToken {
let delegate = GitHubTarballRedirectDelegate(bearerToken: bearerToken)
session = URLSession(configuration: .default, delegate: delegate, delegateQueue: nil)
} else {
session = URLSession.shared
}
return try await session.data(for: request)
}
let maxBytes = Self.maxTarballBytes()
guard data.count <= maxBytes else {
Expand Down Expand Up @@ -76,6 +67,48 @@ struct RepoFetcher {
return RepoFetchOutcome(extractPath: extractPath, tempRoot: tempDir, resolvedCommitSha: fromArchive)
}

/// Downloads a GitHub repository tarball, retrying once without credentials when GitHub rejects
/// an optional bearer token. This lets public repositories remain available after a user's OAuth
/// token expires while private repositories still fail closed on the anonymous retry.
static func downloadTarball(
url: URL,
authToken: String?,
loader: TarballDataLoader
) async throws -> Data {
var request = tarballRequest(url: url, authToken: authToken)
var (data, response) = try await loader(request, authToken)
var status = (response as? HTTPURLResponse)?.statusCode ?? 0

if status == 401, authToken != nil {
request = tarballRequest(url: url, authToken: nil)
do {
(data, response) = try await loader(request, nil)
status = (response as? HTTPURLResponse)?.statusCode ?? 0
guard status == 200 else {
throw RepoFetcherError.fetchFailed(status: 401)
}
} catch {
throw RepoFetcherError.fetchFailed(status: 401)
}
}

guard status == 200 else {
throw RepoFetcherError.fetchFailed(status: status)
}
return data
}

private static func tarballRequest(url: URL, authToken: String?) -> URLRequest {
var request = URLRequest(url: url)
request.httpMethod = "GET"
request.timeoutInterval = fetchTimeoutSeconds()
if let authToken {
request.setValue("Bearer \(authToken)", forHTTPHeaderField: "Authorization")
}
request.setValue("application/vnd.github.v3+json", forHTTPHeaderField: "Accept")
return request
}

/// `GET /repos/{owner}/{repo}/commits/{ref}` — used when the archive folder name does not include a full SHA.
func resolveCommitShaViaApi(owner: String, repo: String, ref: String, token: String?) async throws -> String? {
guard Self.isValidGitHubOwnerOrRepo(owner), Self.isValidGitHubOwnerOrRepo(repo), Self.isValidGitHubRef(ref) else {
Expand Down
15 changes: 15 additions & 0 deletions services/mcp-gateway/Tests/AppTests/MCPAgentVisibilityTests.swift
Original file line number Diff line number Diff line change
Expand Up @@ -71,6 +71,8 @@ struct MCPAgentVisibilityTests {
#expect(resolve.properties?["available_tools"]?.type == "array")
#expect(resolve.properties?["available_tools"]?.items?.type == "object")
#expect(resolve.properties?["context"]?.type == "object")
#expect(resolve.properties?["context"]?.properties?["task"]?.type == "string")
#expect(resolve.properties?["task"]?.type == "string")

let feedback = CapabilitySchemaBuilder.runtimeToolInputSchema(name: "report_skill_feedback")
#expect(Set(feedback.required ?? []) == ["skill_id", "version", "category", "summary", "evidence"])
Expand All @@ -85,6 +87,19 @@ struct MCPAgentVisibilityTests {
}
}

@Test func runtimeContextPreservesNestedAndCompatibilityTaskIdentity() throws {
let nested = try SkillRuntimeToolHandlers.runtimeContext(arguments: [
"context": .object(["task": .string("MCP-10")]),
])
#expect(nested.task == "MCP-10")

let compatibility = try SkillRuntimeToolHandlers.runtimeContext(arguments: [
"context": .object(["task": .string("nested-task")]),
"task": .string("compatibility-task"),
])
#expect(compatibility.task == "compatibility-task")
}

@Test func paginationIsStableAndRejectsWrongContext() throws {
let values = ["a", "b", "c", "d", "e"]
let first = try MCPPaginator.page(values, cursor: nil, scope: "tools:release-a", pageSize: 2)
Expand Down
140 changes: 140 additions & 0 deletions services/mcp-gateway/Tests/AppTests/RepoFetcherTests.swift
Original file line number Diff line number Diff line change
@@ -0,0 +1,140 @@
import Foundation
#if canImport(FoundationNetworking)
import FoundationNetworking
#endif
import Testing
@testable import App

@Suite("Repository fetch authentication fallback")
struct RepoFetcherTests {
@Test("An authenticated 401 retries anonymously and can fetch a public repository")
func retriesPublicRepositoryAnonymously() async throws {
let recorder = TarballRequestRecorder(statuses: [401, 200], successBody: Data("public archive".utf8))
let data = try await RepoFetcher.downloadTarball(
url: try #require(URL(string: "https://api.github.com/repos/example/skills/tarball/main")),
authToken: "stale-token",
loader: { request, token in
try await recorder.load(request, token: token)
}
)

#expect(data == Data("public archive".utf8))
let requests = await recorder.requests
#expect(requests.count == 2)
#expect(requests[0].authorization == "Bearer stale-token")
#expect(requests[0].loaderToken == "stale-token")
#expect(requests[1].authorization == nil)
#expect(requests[1].loaderToken == nil)
}

@Test("A private repository remains inaccessible after the anonymous retry")
func privateRepositoryStillFailsClosed() async throws {
let recorder = TarballRequestRecorder(statuses: [401, 404])

do {
_ = try await RepoFetcher.downloadTarball(
url: try #require(URL(string: "https://api.github.com/repos/example/private/tarball/main")),
authToken: "stale-token",
loader: { request, token in
try await recorder.load(request, token: token)
}
)
Issue.record("Expected the anonymous retry to fail")
} catch RepoFetcherError.fetchFailed(let status) {
#expect(status == 401)
}

let requests = await recorder.requests
#expect(requests.count == 2)
#expect(requests[1].authorization == nil)
}

@Test("A valid credential succeeds without an anonymous retry")
func validCredentialDoesNotRetry() async throws {
let recorder = TarballRequestRecorder(statuses: [200], successBody: Data("private archive".utf8))
let data = try await RepoFetcher.downloadTarball(
url: try #require(URL(string: "https://api.github.com/repos/example/private/tarball/main")),
authToken: "valid-token",
loader: { request, token in
try await recorder.load(request, token: token)
}
)

#expect(data == Data("private archive".utf8))
let requests = await recorder.requests
#expect(requests.count == 1)
#expect(requests[0].authorization == "Bearer valid-token")
}

@Test("Non-authentication failures do not trigger an anonymous retry")
func doesNotRetryOtherFailures() async throws {
let recorder = TarballRequestRecorder(statuses: [403])

do {
_ = try await RepoFetcher.downloadTarball(
url: try #require(URL(string: "https://api.github.com/repos/example/skills/tarball/main")),
authToken: "valid-token",
loader: { request, token in
try await recorder.load(request, token: token)
}
)
Issue.record("Expected the download to fail")
} catch RepoFetcherError.fetchFailed(let status) {
#expect(status == 403)
}

#expect(await recorder.requests.count == 1)
}

@Test("An anonymous 401 is terminal")
func anonymous401IsTerminal() async throws {
let recorder = TarballRequestRecorder(statuses: [401])

do {
_ = try await RepoFetcher.downloadTarball(
url: try #require(URL(string: "https://api.github.com/repos/example/skills/tarball/main")),
authToken: nil,
loader: { request, token in
try await recorder.load(request, token: token)
}
)
Issue.record("Expected the download to fail")
} catch RepoFetcherError.fetchFailed(let status) {
#expect(status == 401)
}

#expect(await recorder.requests.count == 1)
}
}

private actor TarballRequestRecorder {
struct RecordedRequest: Sendable {
let authorization: String?
let loaderToken: String?
}

private var statuses: [Int]
private let successBody: Data
private(set) var requests: [RecordedRequest] = []

init(statuses: [Int], successBody: Data = Data()) {
self.statuses = statuses
self.successBody = successBody
}

func load(_ request: URLRequest, token: String?) throws -> (Data, URLResponse) {
requests.append(RecordedRequest(
authorization: request.value(forHTTPHeaderField: "Authorization"),
loaderToken: token
))
let status = statuses.isEmpty ? 500 : statuses.removeFirst()
let requestURL = try #require(request.url)
let response = try #require(HTTPURLResponse(
url: requestURL,
statusCode: status,
httpVersion: "HTTP/1.1",
headerFields: nil
))
return (status == 200 ? successBody : Data(), response)
}
}
Loading