mirror of
https://github.com/tiennm99/DocsGPT.git
synced 2026-10-11 03:12:55 +00:00
Add a built-in GitHub connector for repository sync and read-only tools
One GitHub connection feeds both a Knowledge source and an agent tool. Users connect with a personal access token, checked against GitHub and named after the account, or, when an admin registers a GitHub App (GITHUB_CLIENT_ID, GITHUB_CLIENT_SECRET, GITHUB_APP_SLUG), with Sign in with GitHub. App tokens expire after eight hours and are refreshed before each sync or tool call; tokens with no expiry are never treated as expired. The catalog gains an optional second sign-in method (oauth_settings, exposed as sign_in_methods) without changing any other connector. Sync lists the repositories the connection can read (the token's own, or the App installations') and ingests one with the connection's token. The tool is GitHub's read-only MCP server (api.githubcopilot.com/mcp/readonly): setup discovers its actions and creates it bound to the connection, and the executor sends the connection's token only to that server.
This commit is contained in:
1 parent
3572d968ec
commit
7b516ade82
15 files changed
+1078
-24
No files matched your search
@@ -1140,7 +1140,25 @@ Confluence Cloud OAuth client secret.
|
||||
|
||||
Type `str`, default unset.
|
||||
|
||||
GitHub PAT with read access to repositories.
|
||||
Instance-wide GitHub token for the public-repository upload. It raises GitHub's rate limit and is never used to read a private repository; users connect their own GitHub account for those.
|
||||
|
||||
### `GITHUB_CLIENT_ID`
|
||||
|
||||
Type `str`, default unset.
|
||||
|
||||
GitHub App client id. With the secret and slug, offers Sign in with GitHub next to tokens.
|
||||
|
||||
### `GITHUB_CLIENT_SECRET`
|
||||
|
||||
Type `str`, default unset.
|
||||
|
||||
GitHub App client secret.
|
||||
|
||||
### `GITHUB_APP_SLUG`
|
||||
|
||||
Type `str`, default unset.
|
||||
|
||||
GitHub App URL name (github.com/apps/<slug>), for the link where users choose repositories.
|
||||
|
||||
### `MCP_OAUTH_REDIRECT_URI`
|
||||
|
||||
|
||||
@@ -1698,8 +1698,13 @@ class ToolExecutor:
|
||||
if tool_data.get("name") == "mcp_tool":
|
||||
from docsgpt.connectors import catalog
|
||||
|
||||
# A connection's secret only goes to the server it was stored for.
|
||||
# A connection's secret only goes to the server it was stored for:
|
||||
# a custom server's own URL, or a built-in connector's MCP server
|
||||
# (GitHub's token only ever goes to GitHub's).
|
||||
stored_for = catalog.base_url(resolved.row.get("server_url"))
|
||||
definition = catalog.get_definition(resolved.connector_key)
|
||||
if not stored_for and definition is not None and definition.publisher == "built_in":
|
||||
stored_for = definition.mcp_base_url or ""
|
||||
if stored_for and stored_for != catalog.base_url(tool_config.get("server_url")):
|
||||
raise service.ConnectionUnavailable(
|
||||
f"{resolved.connector_name or 'This service'} needs to be connected",
|
||||
@@ -1713,8 +1718,10 @@ class ToolExecutor:
|
||||
agent_id=self.agent_id,
|
||||
)
|
||||
tool_config.pop("encrypted_credentials", None)
|
||||
if (resolved.row.get("auth_kind") or "") in ("api_key", "none"):
|
||||
credentials = service.get_credentials(resolved.row)
|
||||
if (resolved.row.get("auth_kind") or "") in ("api_key", "none", "oauth"):
|
||||
# Pasted keys, or the current access token of a built-in OAuth
|
||||
# sign-in (GitHub's App), refreshed first when it has expired.
|
||||
credentials = service.access_credentials(resolved.row)
|
||||
tool_config.update(credentials)
|
||||
tool_config["auth_credentials"] = credentials
|
||||
if tool_data.get("name") == "mcp_tool":
|
||||
|
||||
@@ -91,10 +91,22 @@ class ConnectionsList(Resource):
|
||||
credentials = body.get("credentials")
|
||||
if not isinstance(credentials, dict):
|
||||
return _error("credentials must be an object", 400)
|
||||
label = body.get("label") or None
|
||||
if definition.key == "github":
|
||||
# Check the token now, not at the first sync, and name the
|
||||
# connection after the account rather than a hint of the token.
|
||||
from docsgpt.connectors import github
|
||||
|
||||
try:
|
||||
label = label or github.token_account(credentials.get("access_token"))
|
||||
except github.TokenRejected as err:
|
||||
return _error(str(err), 400, code="invalid_credentials")
|
||||
except service.TransientConnectionError as err:
|
||||
return _error(str(err), 502)
|
||||
try:
|
||||
with db_session() as conn:
|
||||
row, created = service.create_api_key_connection(
|
||||
conn, user_id, definition, credentials, label=(body.get("label") or None),
|
||||
conn, user_id, definition, credentials, label=label,
|
||||
)
|
||||
except service.EncryptionKeyNotConfigured as err:
|
||||
return _error(str(err), 400, code="encryption_key_default")
|
||||
@@ -211,17 +223,38 @@ class ConnectionSetup(Resource):
|
||||
if not user_id:
|
||||
return _unauthorized()
|
||||
body = _json_body()
|
||||
create_tools = body.get("create_tools", True)
|
||||
try:
|
||||
mcp_actions = None
|
||||
with db_readonly() as conn:
|
||||
row = _owned(conn, connection_id, user_id)
|
||||
discover = bool(row) and create_tools and service.needs_mcp_discovery(conn, row)
|
||||
if row is None:
|
||||
return _not_found()
|
||||
if service.normalize_status(row) != service.STATUS_CONNECTED:
|
||||
return _error("Reconnect before setting up", 409, code="reconnect")
|
||||
if discover:
|
||||
# GitHub's tool is its MCP server: read its actions before
|
||||
# the write transaction, not while holding it open.
|
||||
from docsgpt.connectors.mcp import discover_builtin_actions
|
||||
|
||||
try:
|
||||
mcp_actions = discover_builtin_actions(user_id, row)
|
||||
except service.ConnectionUnavailable:
|
||||
return _error("Reconnect before setting up", 409, code="reconnect")
|
||||
except Exception as err:
|
||||
current_app.logger.warning(f"Could not list the MCP server's tools: {err}")
|
||||
return _error("The service's tools could not be reached. Try again.", 502,
|
||||
code="tools_unavailable")
|
||||
with db_session() as conn:
|
||||
row = _owned(conn, connection_id, user_id)
|
||||
if row is None:
|
||||
return _not_found()
|
||||
if service.normalize_status(row) != service.STATUS_CONNECTED:
|
||||
return _error("Reconnect before setting up", 409, code="reconnect")
|
||||
tools = []
|
||||
if body.get("create_tools", True):
|
||||
if create_tools:
|
||||
tools = service.ensure_connection_tools(
|
||||
conn, user_id, row, permissions=body.get("tool_permissions") or None,
|
||||
mcp_actions=mcp_actions,
|
||||
)
|
||||
tool_payload = [service.serialize_tool(tool) for tool in tools]
|
||||
sources = []
|
||||
@@ -258,7 +291,16 @@ def _start_sync(user_id: str, row: dict, sync: dict):
|
||||
frequency = sync.get("frequency") or definition.default_sync_frequency
|
||||
if frequency not in _FREQUENCIES:
|
||||
return ("Unknown sync frequency", 400)
|
||||
name = (sync.get("name") or "").strip() or definition.name
|
||||
name = (sync.get("name") or "").strip()
|
||||
if definition.sync_ingestor == "github":
|
||||
from docsgpt.parser.remote.github_loader import GitHubLoader
|
||||
|
||||
repo = GitHubLoader.normalize_repo(str(items.get("repo_url") or ""))
|
||||
if not repo:
|
||||
return ("Pick a GitHub repository", 400)
|
||||
items = {**items, "repo_url": repo}
|
||||
name = name or repo
|
||||
name = name or definition.name
|
||||
# Validate before claiming the idempotency key: a rejected request must
|
||||
# leave the key free for the corrected retry.
|
||||
if definition.auth_kind == "oauth":
|
||||
@@ -307,6 +349,49 @@ def _start_sync(user_id: str, row: dict, sync: dict):
|
||||
return {"id": source_id, "task_id": task_id or task.id, "name": name, "sync_frequency": frequency}
|
||||
|
||||
|
||||
@connections_ns.route("/connections/<string:connection_id>/repositories")
|
||||
class ConnectionRepositories(Resource):
|
||||
@api.doc(
|
||||
description=(
|
||||
"GitHub: the repositories the connection can read, for the sync picker. "
|
||||
"install_url is where a GitHub App sign-in chooses more repositories."
|
||||
)
|
||||
)
|
||||
def get(self, connection_id: str):
|
||||
from docsgpt.connectors import github
|
||||
|
||||
user_id = _user_id()
|
||||
if not user_id:
|
||||
return _unauthorized()
|
||||
with db_readonly() as conn:
|
||||
row = _owned(conn, connection_id, user_id)
|
||||
if row is None or catalog.connector_key_for_row(row) != "github":
|
||||
return _not_found()
|
||||
app_sign_in = (row.get("auth_kind") or "") == "oauth"
|
||||
try:
|
||||
token = service.access_credentials(row).get("access_token")
|
||||
repositories = github.list_repositories(token or "", app=app_sign_in)
|
||||
except service.ConnectionUnavailable:
|
||||
return _error("Reconnect to continue", 409, code="reconnect")
|
||||
except github.TokenRejected as err:
|
||||
service.mark_reconnect_needed(connection_id, str(err))
|
||||
return _error("Reconnect to continue", 409, code="reconnect")
|
||||
except service.TransientConnectionError:
|
||||
return _error("GitHub is not responding. Try again.", 503)
|
||||
except Exception as err:
|
||||
current_app.logger.error(f"Error listing GitHub repositories: {err}", exc_info=True)
|
||||
return _error("Failed to list repositories", 502)
|
||||
install_url = None
|
||||
if app_sign_in:
|
||||
from docsgpt.core.settings import settings
|
||||
|
||||
slug = settings.GITHUB_APP_SLUG
|
||||
install_url = f"https://github.com/apps/{slug}/installations/new" if slug else None
|
||||
return make_response(
|
||||
jsonify({"success": True, "repositories": repositories, "install_url": install_url}), 200,
|
||||
)
|
||||
|
||||
|
||||
@connections_ns.route("/connections/<string:connection_id>/reconnect")
|
||||
class ConnectionReconnect(Resource):
|
||||
@api.doc(
|
||||
|
||||
@@ -106,7 +106,7 @@ def _render_callback_page(
|
||||
# The script only carries server-side values: the provider key comes from the
|
||||
# supported-connector list rather than the request, and no request text is posted.
|
||||
provider_key = next(
|
||||
(key for key in ConnectorCreator.get_supported_connectors() if key == provider_raw.lower()), None,
|
||||
(key for key in ConnectorCreator.get_auth_providers() if key == provider_raw.lower()), None,
|
||||
)
|
||||
payload = None
|
||||
if provider_key and status == "success" and session_token:
|
||||
@@ -177,12 +177,24 @@ def _render_callback_page(
|
||||
|
||||
|
||||
|
||||
def build_authorization(provider: str, user_id: str, connection_id: Optional[str] = None) -> dict:
|
||||
def build_authorization(
|
||||
provider: str, user_id: str, connection_id: Optional[str] = None, *, install: bool = False,
|
||||
) -> dict:
|
||||
"""Start an OAuth sign-in for ``provider`` and return its authorization URL.
|
||||
|
||||
Args:
|
||||
provider: The connector, e.g. ``google_drive`` or ``github``.
|
||||
user_id: The caller.
|
||||
connection_id: The caller's connection to sign in again, if any.
|
||||
install: GitHub: send the user to install the GitHub App (where the
|
||||
repositories it can read are chosen) instead of straight to
|
||||
authorization. GitHub returns to the callback with the same
|
||||
state when the app requests authorization during installation.
|
||||
|
||||
Raises:
|
||||
service.EncryptionKeyNotConfigured: See ``ensure_can_store_credentials``.
|
||||
service.ConnectionUnavailable: ``connection_id`` is not the caller's.
|
||||
ValueError: ``install`` for a provider with no installation page.
|
||||
"""
|
||||
service.ensure_can_store_credentials()
|
||||
with db_session() as conn:
|
||||
@@ -191,8 +203,11 @@ def build_authorization(provider: str, user_id: str, connection_id: Optional[str
|
||||
json.dumps({"provider": provider, "object_id": str(session_row["id"])}).encode()
|
||||
).decode()
|
||||
auth = ConnectorCreator.create_auth(provider)
|
||||
if install and not hasattr(auth, "get_installation_url"):
|
||||
raise ValueError(f"{provider} has no installation page")
|
||||
url = auth.get_installation_url(state=state) if install else auth.get_authorization_url(state=state)
|
||||
return {
|
||||
"authorization_url": auth.get_authorization_url(state=state),
|
||||
"authorization_url": url,
|
||||
"state": state,
|
||||
"callback_origin": _origin_of(settings.CONNECTOR_REDIRECT_BASE_URI),
|
||||
}
|
||||
@@ -207,7 +222,7 @@ class ConnectorAuth(Resource):
|
||||
if not provider:
|
||||
return make_response(jsonify({"success": False, "error": "Missing provider"}), 400)
|
||||
|
||||
if not ConnectorCreator.is_supported(provider):
|
||||
if not (ConnectorCreator.is_supported(provider) or ConnectorCreator.has_auth(provider)):
|
||||
return make_response(jsonify({"success": False, "error": f"Unsupported provider: {provider}"}), 400)
|
||||
|
||||
decoded_token = request.decoded_token
|
||||
@@ -215,7 +230,10 @@ class ConnectorAuth(Resource):
|
||||
return make_response(jsonify({"success": False, "error": "Unauthorized"}), 401)
|
||||
user_id = decoded_token.get('sub')
|
||||
try:
|
||||
started = build_authorization(provider, user_id, request.args.get("connection_id") or None)
|
||||
started = build_authorization(
|
||||
provider, user_id, request.args.get("connection_id") or None,
|
||||
install=request.args.get("install") in ("1", "true"),
|
||||
)
|
||||
except service.EncryptionKeyNotConfigured as err:
|
||||
return make_response(
|
||||
jsonify({"success": False, "error": str(err), "code": "encryption_key_default"}), 400,
|
||||
@@ -251,12 +269,24 @@ class ConnectorsCallback(Resource):
|
||||
state = request.args.get('state')
|
||||
error = request.args.get('error')
|
||||
|
||||
if not state and request.args.get('installation_id'):
|
||||
# The GitHub App was installed from GitHub itself, not from a
|
||||
# DocsGPT sign-in: there is no state to tie the code to a user,
|
||||
# so the code is ignored and the user goes back to DocsGPT.
|
||||
return _render_callback_page(
|
||||
"success",
|
||||
"The GitHub App is installed. Return to DocsGPT and refresh the repository list.",
|
||||
"github",
|
||||
)
|
||||
|
||||
state_dict = json.loads(base64.urlsafe_b64decode(state.encode()).decode())
|
||||
provider = state_dict.get("provider")
|
||||
state_object_id = state_dict.get("object_id")
|
||||
|
||||
# Validate provider
|
||||
if not provider or not isinstance(provider, str) or not ConnectorCreator.is_supported(provider):
|
||||
if not provider or not isinstance(provider, str) or not (
|
||||
ConnectorCreator.is_supported(provider) or ConnectorCreator.has_auth(provider)
|
||||
):
|
||||
return redirect(build_callback_redirect({
|
||||
"status": "error",
|
||||
"message": "Invalid provider"
|
||||
@@ -295,7 +325,9 @@ class ConnectorsCallback(Resource):
|
||||
user_info = drive_service.about().get(fields="user").execute()
|
||||
user_email = user_info.get('user', {}).get('emailAddress', 'Connected User')
|
||||
else:
|
||||
user_email = token_info.get('user_info', {}).get('email', 'Connected User')
|
||||
# GitHub names the account by its login, the others by email.
|
||||
user_info = token_info.get('user_info') or {}
|
||||
user_email = user_info.get('email') or user_info.get('login') or 'Connected User'
|
||||
|
||||
except Exception as e:
|
||||
current_app.logger.warning(f"Could not get user info: {e}")
|
||||
|
||||
@@ -72,12 +72,17 @@ class ConnectorDefinition:
|
||||
tool_templates: ``user_tools`` names created on connect.
|
||||
setup: What the wizard does after sign-in: ``tools`` is ``auto``
|
||||
(created and enabled), ``ask`` or ``off``; ``sync`` likewise.
|
||||
mcp_url: MCP endpoint for presets.
|
||||
mcp_url: MCP endpoint for presets, and for a built-in connector whose
|
||||
tool is its service's own MCP server (GitHub).
|
||||
publisher: ``built_in``, ``preset`` or ``custom``.
|
||||
docs_url: Setup guide for admins.
|
||||
oauth_scopes: Scopes an MCP preset requests.
|
||||
part_of: Another connector this one is shown under, for one service
|
||||
offered two ways (the Atlassian MCP preset under Confluence).
|
||||
oauth_settings: Server settings that add a second sign-in, OAuth,
|
||||
to an ``api_key`` connector (GitHub's "Sign in with GitHub"
|
||||
through a GitHub App). Unlike ``required_settings`` the
|
||||
connector works without them, with pasted credentials only.
|
||||
"""
|
||||
|
||||
key: str
|
||||
@@ -99,6 +104,17 @@ class ConnectorDefinition:
|
||||
docs_url: Optional[str] = None
|
||||
oauth_scopes: tuple[str, ...] = ()
|
||||
part_of: Optional[str] = None
|
||||
oauth_settings: tuple[str, ...] = ()
|
||||
|
||||
@property
|
||||
def oauth_configured(self) -> bool:
|
||||
"""Whether the optional OAuth sign-in has every server setting it needs."""
|
||||
return bool(self.oauth_settings) and all(getattr(settings, name, None) for name in self.oauth_settings)
|
||||
|
||||
@property
|
||||
def sign_in_methods(self) -> list[str]:
|
||||
"""How a user can connect, preferred first: ``auth_kind``, after OAuth when that is set up."""
|
||||
return ["oauth", self.auth_kind] if self.oauth_configured else [self.auth_kind]
|
||||
|
||||
@property
|
||||
def missing_settings(self) -> list[str]:
|
||||
@@ -136,6 +152,7 @@ class ConnectorDefinition:
|
||||
"docs_url": self.docs_url,
|
||||
"oauth_scopes": list(self.oauth_scopes),
|
||||
"part_of": self.part_of,
|
||||
"sign_in_methods": self.sign_in_methods,
|
||||
}
|
||||
|
||||
|
||||
@@ -148,6 +165,11 @@ def base_url(url: Optional[str]) -> str:
|
||||
|
||||
|
||||
_DOCS = "https://docs.docsgpt.cloud/Guides/Connectors"
|
||||
GITHUB_MCP_URL = "https://api.githubcopilot.com/mcp/readonly"
|
||||
|
||||
# Tool templates any connector's server can use (an MCP server, an OpenAPI
|
||||
# spec): a tool made from one belongs to its connection, not to a connector.
|
||||
_GENERIC_TOOL_TEMPLATES = frozenset({"mcp_tool", "api_tool"})
|
||||
|
||||
_BUILT_IN: tuple[ConnectorDefinition, ...] = (
|
||||
ConnectorDefinition(
|
||||
@@ -189,6 +211,25 @@ _BUILT_IN: tuple[ConnectorDefinition, ...] = (
|
||||
setup={"tools": "off", "sync": "ask"},
|
||||
docs_url=f"{_DOCS}#confluence",
|
||||
),
|
||||
ConnectorDefinition(
|
||||
key="github",
|
||||
name="GitHub",
|
||||
description="Sync repositories into Knowledge and let agents read code, issues and pull requests.",
|
||||
icon="github",
|
||||
category="dev",
|
||||
# A token works with no admin setup; a GitHub App adds Sign in with GitHub.
|
||||
auth_kind="api_key",
|
||||
capabilities=("sync", "read"),
|
||||
credential_fields=(CredentialField("access_token", "Personal access token"),),
|
||||
oauth_settings=("GITHUB_CLIENT_ID", "GITHUB_CLIENT_SECRET", "GITHUB_APP_SLUG"),
|
||||
sync_ingestor="github",
|
||||
setup_fields=(CredentialField("repo_url", "Repository", secret=False),),
|
||||
# GitHub's own MCP server, read-only: agents look things up, never change them.
|
||||
tool_templates=("mcp_tool",),
|
||||
mcp_url=GITHUB_MCP_URL,
|
||||
setup={"tools": "ask", "sync": "ask"},
|
||||
docs_url=f"{_DOCS}#github",
|
||||
),
|
||||
ConnectorDefinition(
|
||||
key="s3",
|
||||
name="Amazon S3",
|
||||
@@ -373,6 +414,8 @@ def preset_for_url(url: Optional[str]) -> Optional[ConnectorDefinition]:
|
||||
|
||||
def definition_for_tool(tool_name: str) -> Optional[ConnectorDefinition]:
|
||||
"""The built-in connector that provides the ``user_tools`` template ``tool_name``."""
|
||||
if tool_name in _GENERIC_TOOL_TEMPLATES:
|
||||
return None
|
||||
for definition in _BUILT_IN:
|
||||
if definition.publisher == "built_in" and tool_name in definition.tool_templates:
|
||||
return definition
|
||||
@@ -407,6 +450,7 @@ def tool_connector_keys() -> set[str]:
|
||||
for definition in _BUILT_IN
|
||||
if definition.publisher == "built_in"
|
||||
for name in definition.tool_templates
|
||||
if name not in _GENERIC_TOOL_TEMPLATES
|
||||
}
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,132 @@
|
||||
"""GitHub account lookups for the GitHub connector: who a token is, what it can read.
|
||||
|
||||
A connection signs in with a personal access token (``api_key``) or through
|
||||
the GitHub App (``oauth``). A token lists the repositories it was granted;
|
||||
an App sign-in lists the repositories of the App's installations the user
|
||||
can see, which the user chooses on GitHub when installing the App.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Optional
|
||||
|
||||
import requests
|
||||
|
||||
from docsgpt.connectors.service import TransientConnectionError
|
||||
|
||||
API_URL = "https://api.github.com"
|
||||
_PAGE_SIZE = 100
|
||||
# A picker, not an export: enough for any one person's list.
|
||||
MAX_REPOSITORIES = 1000
|
||||
|
||||
|
||||
class TokenRejected(ValueError):
|
||||
"""GitHub answered 401: the token is wrong, expired or revoked."""
|
||||
|
||||
|
||||
def _headers(token: str) -> dict:
|
||||
return {
|
||||
"Authorization": f"Bearer {token}",
|
||||
"Accept": "application/vnd.github+json",
|
||||
"X-GitHub-Api-Version": "2022-11-28",
|
||||
}
|
||||
|
||||
|
||||
def _get(url: str, token: str, params: Optional[dict] = None):
|
||||
"""GET ``url`` as ``token`` and return its JSON.
|
||||
|
||||
Raises:
|
||||
TokenRejected: 401.
|
||||
TransientConnectionError: Network trouble, rate limiting or a 5xx.
|
||||
ValueError: Any other refusal.
|
||||
"""
|
||||
try:
|
||||
response = requests.get(url, headers=_headers(token), params=params, timeout=30)
|
||||
except requests.RequestException as exc:
|
||||
raise TransientConnectionError(f"GitHub did not answer: {type(exc).__name__}") from exc
|
||||
if response.status_code == 401:
|
||||
raise TokenRejected("GitHub did not accept this token.")
|
||||
if response.status_code == 429 or response.status_code >= 500 or (
|
||||
response.status_code == 403 and response.headers.get("X-RateLimit-Remaining") == "0"
|
||||
):
|
||||
raise TransientConnectionError(f"GitHub is busy ({response.status_code}). Try again.")
|
||||
if response.status_code >= 400:
|
||||
raise ValueError(f"GitHub refused the request ({response.status_code}).")
|
||||
return response.json()
|
||||
|
||||
|
||||
def token_account(token: Optional[str]) -> str:
|
||||
"""The login of the account ``token`` belongs to.
|
||||
|
||||
Raises:
|
||||
TokenRejected: No token, or GitHub did not accept it.
|
||||
TransientConnectionError: GitHub could not be reached.
|
||||
"""
|
||||
if not token or not str(token).strip():
|
||||
raise TokenRejected("Paste a GitHub token.")
|
||||
user = _get(f"{API_URL}/user", str(token).strip())
|
||||
login = user.get("login") if isinstance(user, dict) else None
|
||||
if not login:
|
||||
raise TokenRejected("GitHub did not say whose token this is.")
|
||||
return str(login)
|
||||
|
||||
|
||||
def _summary(repo: dict) -> dict:
|
||||
return {
|
||||
"full_name": repo.get("full_name"),
|
||||
"private": bool(repo.get("private")),
|
||||
"description": repo.get("description") or "",
|
||||
"default_branch": repo.get("default_branch") or "",
|
||||
"updated_at": repo.get("pushed_at") or repo.get("updated_at"),
|
||||
"html_url": repo.get("html_url") or "",
|
||||
}
|
||||
|
||||
|
||||
def _paged(url: str, token: str, key: Optional[str] = None, params: Optional[dict] = None):
|
||||
"""Every item of a paginated list endpoint, up to :data:`MAX_REPOSITORIES`."""
|
||||
page = 1
|
||||
fetched = 0
|
||||
while fetched < MAX_REPOSITORIES:
|
||||
payload = _get(url, token, {**(params or {}), "per_page": _PAGE_SIZE, "page": page})
|
||||
items = payload.get(key, []) if key else payload
|
||||
if not isinstance(items, list):
|
||||
return
|
||||
for item in items:
|
||||
if isinstance(item, dict):
|
||||
fetched += 1
|
||||
yield item
|
||||
if len(items) < _PAGE_SIZE:
|
||||
return
|
||||
page += 1
|
||||
|
||||
|
||||
def list_repositories(token: str, *, app: bool) -> list[dict]:
|
||||
"""Repositories ``token`` can read, most recently pushed first.
|
||||
|
||||
Args:
|
||||
token: The connection's access token.
|
||||
app: Whether it is a GitHub App user token, whose repositories are
|
||||
those of the App's installations.
|
||||
|
||||
Returns:
|
||||
``{full_name, private, description, default_branch, updated_at,
|
||||
html_url}`` per repository.
|
||||
|
||||
Raises:
|
||||
TokenRejected: GitHub did not accept the token.
|
||||
TransientConnectionError: GitHub could not be reached.
|
||||
"""
|
||||
if app:
|
||||
raw = []
|
||||
for installation in _paged(f"{API_URL}/user/installations", token, "installations"):
|
||||
raw.extend(_paged(
|
||||
f"{API_URL}/user/installations/{installation.get('id')}/repositories", token, "repositories",
|
||||
))
|
||||
else:
|
||||
raw = list(_paged(f"{API_URL}/user/repos", token, params={"sort": "pushed"}))
|
||||
seen: dict[str, dict] = {}
|
||||
for repo in raw:
|
||||
name = repo.get("full_name")
|
||||
if name and name not in seen:
|
||||
seen[name] = _summary(repo)
|
||||
return sorted(seen.values(), key=lambda r: r["updated_at"] or "", reverse=True)[:MAX_REPOSITORIES]
|
||||
@@ -2,7 +2,7 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from docsgpt.connectors import service
|
||||
from docsgpt.connectors import catalog, service
|
||||
from docsgpt.connectors.permissions import apply_default_permissions
|
||||
from docsgpt.storage.db.repositories.connector_sessions import ConnectorSessionsRepository
|
||||
from docsgpt.storage.db.repositories.user_tools import UserToolsRepository
|
||||
@@ -13,8 +13,9 @@ def _discover(user_id: str, connection: dict, tool: dict) -> list[dict]:
|
||||
from docsgpt.agents.tools.mcp_tool import MCPTool
|
||||
|
||||
config = {k: v for k, v in (tool.get("config") or {}).items() if k != "encrypted_credentials"}
|
||||
if (connection.get("auth_kind") or "") in ("api_key", "none"):
|
||||
config["auth_credentials"] = service.get_credentials(connection)
|
||||
if (connection.get("auth_kind") or "") in ("api_key", "none", "oauth"):
|
||||
# Pasted keys, or the access token of a built-in OAuth sign-in (GitHub's App).
|
||||
config["auth_credentials"] = service.access_credentials(connection)
|
||||
elif service.normalize_status(connection) != service.STATUS_CONNECTED:
|
||||
raise service.ConnectionUnavailable("Reconnect to continue", connection_id=str(connection["id"]))
|
||||
config["connection_id"] = str(connection["id"])
|
||||
@@ -24,6 +25,21 @@ def _discover(user_id: str, connection: dict, tool: dict) -> list[dict]:
|
||||
return mcp_tool.get_actions_metadata()
|
||||
|
||||
|
||||
def discover_builtin_actions(user_id: str, connection: dict) -> list[dict]:
|
||||
"""The actions of a built-in connector's MCP server (GitHub's), read with the connection.
|
||||
|
||||
Raises:
|
||||
ValueError: The connector has no MCP server of its own.
|
||||
service.ConnectionUnavailable: The connection needs reconnecting.
|
||||
Exception: The server could not be reached or listed.
|
||||
"""
|
||||
definition = catalog.get_definition(catalog.connector_key_for_row(connection))
|
||||
config = service.builtin_mcp_config(definition)
|
||||
if config is None:
|
||||
raise ValueError("This connector has no MCP server")
|
||||
return _discover(user_id, connection, {"config": config})
|
||||
|
||||
|
||||
def refresh_mcp_tools(user_id: str, connection: dict) -> dict:
|
||||
"""Re-read the actions of every MCP tool on a connection.
|
||||
|
||||
|
||||
@@ -1123,11 +1123,54 @@ def split_secrets(config: dict, config_requirements: dict) -> tuple[dict, dict]:
|
||||
return public, secrets
|
||||
|
||||
|
||||
def ensure_connection_tools(conn, user_id: str, connection: dict, permissions: Optional[dict] = None) -> list[dict]:
|
||||
def builtin_mcp_config(definition: Optional[ConnectorDefinition]) -> Optional[dict]:
|
||||
"""Tool config of a built-in connector whose tool is its service's MCP server.
|
||||
|
||||
The token is not part of it: the executor reads it from the connection
|
||||
at run time and sends it as a bearer token.
|
||||
|
||||
Returns:
|
||||
The config, or None when the connector has no such tool.
|
||||
"""
|
||||
if (
|
||||
definition is None
|
||||
or definition.publisher != "built_in"
|
||||
or not definition.mcp_url
|
||||
or "mcp_tool" not in definition.tool_templates
|
||||
):
|
||||
return None
|
||||
return {"server_url": definition.mcp_url, "auth_type": "bearer", "transport_type": "http", "timeout": 30}
|
||||
|
||||
|
||||
def needs_mcp_discovery(conn, connection: dict) -> bool:
|
||||
"""Whether creating this connection's tools first needs its MCP server's actions."""
|
||||
definition = catalog.get_definition(catalog.connector_key_for_row(connection))
|
||||
if builtin_mcp_config(definition) is None:
|
||||
return False
|
||||
have = {tool.get("name") for tool in ConnectorSessionsRepository(conn).list_tools(str(connection["id"]))}
|
||||
return "mcp_tool" not in have
|
||||
|
||||
|
||||
def ensure_connection_tools(
|
||||
conn,
|
||||
user_id: str,
|
||||
connection: dict,
|
||||
permissions: Optional[dict] = None,
|
||||
mcp_actions: Optional[list] = None,
|
||||
) -> list[dict]:
|
||||
"""Create the connector's tools once; later calls return the existing ones.
|
||||
|
||||
This is what makes the setup step idempotent for tools: a retried or
|
||||
repeated setup never creates a second Telegram tool for the same bot.
|
||||
|
||||
Args:
|
||||
conn: Open connection inside a transaction.
|
||||
user_id: The owner.
|
||||
connection: The connection row.
|
||||
permissions: Per-action permission overrides.
|
||||
mcp_actions: The actions of a built-in connector's MCP server
|
||||
(GitHub's), discovered by the caller outside this transaction.
|
||||
Without them that tool is not created.
|
||||
"""
|
||||
repo = ConnectorSessionsRepository(conn)
|
||||
existing = repo.list_tools(str(connection["id"]))
|
||||
@@ -1136,10 +1179,19 @@ def ensure_connection_tools(conn, user_id: str, connection: dict, permissions: O
|
||||
if not definition or not definition.tool_templates:
|
||||
return existing
|
||||
have = {tool.get("name") for tool in existing}
|
||||
mcp_config = builtin_mcp_config(definition)
|
||||
created = []
|
||||
for template in definition.tool_templates:
|
||||
if template in have or template in ("mcp_tool", "api_tool"):
|
||||
# MCP tools are created by the MCP save flow, which has the
|
||||
if template in have:
|
||||
continue
|
||||
if template == "mcp_tool" and mcp_config is not None and mcp_actions is not None:
|
||||
created.append(create_tool_for_connection(
|
||||
conn, user_id, connection, template=template, config=mcp_config,
|
||||
actions=mcp_actions, permissions=permissions,
|
||||
))
|
||||
continue
|
||||
if template in ("mcp_tool", "api_tool"):
|
||||
# Other MCP tools are created by the MCP save flow, which has the
|
||||
# discovered actions; OpenAPI tools come from an imported spec.
|
||||
continue
|
||||
created.append(
|
||||
|
||||
@@ -44,7 +44,23 @@ class ConnectorSettings(SettingsGroup):
|
||||
CONFLUENCE_CLIENT_SECRET: Optional[str] = Field(default=None, description="Confluence Cloud OAuth client secret.")
|
||||
|
||||
# GitHub source.
|
||||
GITHUB_ACCESS_TOKEN: Optional[str] = Field(default=None, description="GitHub PAT with read access to repositories.")
|
||||
GITHUB_ACCESS_TOKEN: Optional[str] = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"Instance-wide GitHub token for the public-repository upload. It raises GitHub's rate limit and is never "
|
||||
"used to read a private repository; users connect their own GitHub account for those."
|
||||
),
|
||||
)
|
||||
# GitHub App behind "Sign in with GitHub" on the GitHub connector.
|
||||
GITHUB_CLIENT_ID: Optional[str] = Field(
|
||||
default=None,
|
||||
description="GitHub App client id. With the secret and slug, offers Sign in with GitHub next to tokens.",
|
||||
)
|
||||
GITHUB_CLIENT_SECRET: Optional[str] = Field(default=None, description="GitHub App client secret.")
|
||||
GITHUB_APP_SLUG: Optional[str] = Field(
|
||||
default=None,
|
||||
description="GitHub App URL name (github.com/apps/<slug>), for the link where users choose repositories.",
|
||||
)
|
||||
|
||||
MCP_OAUTH_REDIRECT_URI: Optional[str] = Field(
|
||||
default=None, description="Public callback URL for MCP OAuth; unset derives it from CONNECTOR_REDIRECT_BASE_URI."
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
from docsgpt.parser.connectors.confluence.auth import ConfluenceAuth
|
||||
from docsgpt.parser.connectors.confluence.loader import ConfluenceLoader
|
||||
from docsgpt.parser.connectors.github.auth import GitHubAuth
|
||||
from docsgpt.parser.connectors.google_drive.auth import GoogleDriveAuth
|
||||
from docsgpt.parser.connectors.google_drive.loader import GoogleDriveLoader
|
||||
from docsgpt.parser.connectors.share_point.auth import SharePointAuth
|
||||
@@ -20,8 +21,11 @@ class ConnectorCreator:
|
||||
"share_point": SharePointLoader,
|
||||
}
|
||||
|
||||
# GitHub signs in here too, but its repositories are read by the remote
|
||||
# GitHub loader, so it has an auth provider and no connector class.
|
||||
auth_providers = {
|
||||
"confluence": ConfluenceAuth,
|
||||
"github": GitHubAuth,
|
||||
"google_drive": GoogleDriveAuth,
|
||||
"share_point": SharePointAuth,
|
||||
}
|
||||
@@ -75,6 +79,23 @@ class ConnectorCreator:
|
||||
"""
|
||||
return list(cls.connectors.keys())
|
||||
|
||||
@classmethod
|
||||
def has_auth(cls, connector_type: str) -> bool:
|
||||
"""Whether ``connector_type`` signs in through the OAuth callback.
|
||||
|
||||
Args:
|
||||
connector_type: Provider key, e.g. ``google_drive`` or ``github``.
|
||||
|
||||
Returns:
|
||||
True when an auth provider is registered for it.
|
||||
"""
|
||||
return (connector_type or "").lower() in cls.auth_providers
|
||||
|
||||
@classmethod
|
||||
def get_auth_providers(cls) -> list:
|
||||
"""Provider keys that sign in through the OAuth callback."""
|
||||
return list(cls.auth_providers.keys())
|
||||
|
||||
@classmethod
|
||||
def is_supported(cls, connector_type):
|
||||
"""
|
||||
|
||||
@@ -0,0 +1,9 @@
|
||||
"""GitHub App user sign-in for the GitHub connector.
|
||||
|
||||
Repositories are read by the remote GitHub loader
|
||||
(``docsgpt.parser.remote.github_loader``); this package only signs users in.
|
||||
"""
|
||||
|
||||
from .auth import GitHubAuth
|
||||
|
||||
__all__ = ["GitHubAuth"]
|
||||
@@ -0,0 +1,148 @@
|
||||
"""Sign in with GitHub through a GitHub App (user access tokens).
|
||||
|
||||
A GitHub App's user token can read what both the user and the app's
|
||||
installations can see: the repositories are chosen when the user installs
|
||||
the app, not with OAuth scopes. Tokens last eight hours and come with a
|
||||
six-month refresh token, unless the app turned token expiry off, in which
|
||||
case they carry no expiry and never need refreshing.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import datetime
|
||||
import logging
|
||||
from typing import Any, Dict, Optional
|
||||
from urllib.parse import urlencode
|
||||
|
||||
import requests
|
||||
|
||||
from docsgpt.core.settings import settings
|
||||
from docsgpt.parser.connectors.base import BaseConnectorAuth
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
API_URL = "https://api.github.com"
|
||||
|
||||
|
||||
class GitHubAuth(BaseConnectorAuth):
|
||||
"""OAuth web flow of the GitHub App set in ``GITHUB_CLIENT_ID`` and friends."""
|
||||
|
||||
AUTH_URL = "https://github.com/login/oauth/authorize"
|
||||
TOKEN_URL = "https://github.com/login/oauth/access_token"
|
||||
# Refresh this long before the stated expiry, so a token never lapses mid-request.
|
||||
EXPIRY_MARGIN = datetime.timedelta(minutes=5)
|
||||
|
||||
def __init__(self):
|
||||
self.client_id = settings.GITHUB_CLIENT_ID
|
||||
self.client_secret = settings.GITHUB_CLIENT_SECRET
|
||||
self.app_slug = settings.GITHUB_APP_SLUG
|
||||
self.redirect_uri = settings.CONNECTOR_REDIRECT_BASE_URI
|
||||
if not self.client_id or not self.client_secret:
|
||||
raise ValueError(
|
||||
"GitHub App credentials not configured. "
|
||||
"Please set GITHUB_CLIENT_ID and GITHUB_CLIENT_SECRET in settings."
|
||||
)
|
||||
|
||||
def get_authorization_url(self, state: Optional[str] = None) -> str:
|
||||
"""The GitHub page that asks the user to authorize the app."""
|
||||
params = {"client_id": self.client_id, "redirect_uri": self.redirect_uri, "state": state}
|
||||
return f"{self.AUTH_URL}?{urlencode({k: v for k, v in params.items() if v})}"
|
||||
|
||||
def get_installation_url(self, state: Optional[str] = None) -> str:
|
||||
"""The GitHub page where the user installs the app and picks repositories.
|
||||
|
||||
With "Request user authorization (OAuth) during installation" on, GitHub
|
||||
sends the user back to the callback with a code and this ``state``.
|
||||
"""
|
||||
base = f"https://github.com/apps/{self.app_slug}/installations/new"
|
||||
return f"{base}?{urlencode({'state': state})}" if state else base
|
||||
|
||||
def _token_request(self, data: Dict[str, str]) -> Dict[str, Any]:
|
||||
"""POST to the token endpoint; GitHub reports failures as 200 with ``error``."""
|
||||
response = requests.post(
|
||||
self.TOKEN_URL,
|
||||
data={"client_id": self.client_id, "client_secret": self.client_secret, **data},
|
||||
headers={"Accept": "application/json"},
|
||||
timeout=30,
|
||||
)
|
||||
response.raise_for_status()
|
||||
payload = response.json()
|
||||
if not isinstance(payload, dict) or payload.get("error") or not payload.get("access_token"):
|
||||
error = payload.get("error") if isinstance(payload, dict) else None
|
||||
raise ValueError(f"GitHub refused the sign-in: {error or 'no access token returned'}")
|
||||
return payload
|
||||
|
||||
@staticmethod
|
||||
def _tokens(payload: Dict[str, Any], refresh_token: Optional[str] = None) -> Dict[str, Any]:
|
||||
"""Token info from a token response; ``expiry`` is None for non-expiring tokens."""
|
||||
expires_in = payload.get("expires_in")
|
||||
expiry = None
|
||||
if expires_in:
|
||||
expiry = (
|
||||
datetime.datetime.now(datetime.timezone.utc) + datetime.timedelta(seconds=int(expires_in))
|
||||
).isoformat()
|
||||
return {
|
||||
"access_token": payload["access_token"],
|
||||
"refresh_token": payload.get("refresh_token") or refresh_token,
|
||||
"token_uri": GitHubAuth.TOKEN_URL,
|
||||
"expiry": expiry,
|
||||
}
|
||||
|
||||
def exchange_code_for_tokens(self, authorization_code: str) -> Dict[str, Any]:
|
||||
"""Trade the callback's code for tokens, plus the account's login.
|
||||
|
||||
Raises:
|
||||
ValueError: No code, or GitHub refused it.
|
||||
"""
|
||||
if not authorization_code:
|
||||
raise ValueError("Authorization code is required")
|
||||
payload = self._token_request({"code": authorization_code, "redirect_uri": self.redirect_uri})
|
||||
token_info = self._tokens(payload)
|
||||
token_info["user_info"] = self._fetch_user(token_info["access_token"])
|
||||
return token_info
|
||||
|
||||
def refresh_access_token(self, refresh_token: str) -> Dict[str, Any]:
|
||||
"""A new access token (and a new refresh token: GitHub rotates them).
|
||||
|
||||
Raises:
|
||||
ValueError: The refresh token was refused (expired or revoked).
|
||||
"""
|
||||
if not refresh_token:
|
||||
raise ValueError("Refresh token is required")
|
||||
payload = self._token_request({"grant_type": "refresh_token", "refresh_token": refresh_token})
|
||||
return self._tokens(payload, refresh_token)
|
||||
|
||||
def is_token_expired(self, token_info: Dict[str, Any]) -> bool:
|
||||
"""Whether the access token is (about to be) expired.
|
||||
|
||||
A token with no expiry never expires: the app has token expiry
|
||||
turned off.
|
||||
"""
|
||||
if not token_info or not token_info.get("access_token"):
|
||||
return True
|
||||
expiry = token_info.get("expiry")
|
||||
if not expiry:
|
||||
return False
|
||||
try:
|
||||
expiry_dt = datetime.datetime.fromisoformat(expiry)
|
||||
except (TypeError, ValueError):
|
||||
return True
|
||||
if expiry_dt.tzinfo is None:
|
||||
expiry_dt = expiry_dt.replace(tzinfo=datetime.timezone.utc)
|
||||
return datetime.datetime.now(datetime.timezone.utc) >= expiry_dt - self.EXPIRY_MARGIN
|
||||
|
||||
@staticmethod
|
||||
def _fetch_user(access_token: str) -> Dict[str, Any]:
|
||||
"""``{login, name}`` of the signed-in account; empty when GitHub does not say."""
|
||||
try:
|
||||
response = requests.get(
|
||||
f"{API_URL}/user",
|
||||
headers={"Authorization": f"Bearer {access_token}", "Accept": "application/vnd.github+json"},
|
||||
timeout=30,
|
||||
)
|
||||
response.raise_for_status()
|
||||
user = response.json()
|
||||
except Exception as exc: # the sign-in still works; only its label is missing
|
||||
logger.warning("Could not read the GitHub account: %s", type(exc).__name__)
|
||||
return {}
|
||||
return {"login": user.get("login", ""), "name": user.get("name") or ""}
|
||||
@@ -39,6 +39,9 @@ class TestDefinitions:
|
||||
if definition.publisher != "built_in":
|
||||
continue
|
||||
for tool_name in definition.tool_templates:
|
||||
if tool_name == "mcp_tool":
|
||||
# GitHub's MCP server gets the connection's token as a bearer token.
|
||||
continue
|
||||
requirements = tools[tool_name].get_config_requirements()
|
||||
secret_keys = {k for k, spec in requirements.items() if spec.get("secret")}
|
||||
field_keys = {f.key for f in definition.credential_fields}
|
||||
|
||||
@@ -0,0 +1,343 @@
|
||||
"""The built-in GitHub connector: two sign-ins, repository sync and read-only MCP tools."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from contextlib import contextmanager
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from flask import Flask
|
||||
from sqlalchemy import text
|
||||
|
||||
from docsgpt.connectors import catalog, service
|
||||
from docsgpt.security.encryption import encrypt_json
|
||||
|
||||
READONLY_MCP = "https://api.githubcopilot.com/mcp/readonly"
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def app():
|
||||
return Flask(__name__)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def app_settings(monkeypatch):
|
||||
from docsgpt.core.settings import settings
|
||||
|
||||
monkeypatch.setattr(settings, "GITHUB_CLIENT_ID", "Iv1.client")
|
||||
monkeypatch.setattr(settings, "GITHUB_CLIENT_SECRET", "app-secret")
|
||||
monkeypatch.setattr(settings, "GITHUB_APP_SLUG", "docsgpt-acme")
|
||||
|
||||
|
||||
@contextmanager
|
||||
def _db(conn):
|
||||
@contextmanager
|
||||
def _yield():
|
||||
yield conn
|
||||
|
||||
with patch.multiple("docsgpt.api.connector.connections", db_session=_yield, db_readonly=_yield), \
|
||||
patch.multiple("docsgpt.connectors.service", db_session=_yield, db_readonly=_yield), \
|
||||
patch.multiple("docsgpt.connectors.mcp", db_session=_yield, db_readonly=_yield), \
|
||||
patch.multiple("docsgpt.connectors.resolve", db_readonly=_yield), \
|
||||
patch.multiple("docsgpt.api.connector.routes", db_session=_yield, db_readonly=_yield):
|
||||
yield
|
||||
|
||||
|
||||
def _call(app, resource, method, path, user="alice", body=None, args=(), query=None):
|
||||
with app.test_request_context(path, method=method.upper(), json=body, query_string=query):
|
||||
from flask import request
|
||||
|
||||
request.decoded_token = {"sub": user} if user else None
|
||||
return getattr(resource(), method)(*args)
|
||||
|
||||
|
||||
def _connection(conn, user="alice", auth_kind="api_key", secrets=None, status="connected", label="octocat") -> str:
|
||||
secrets = secrets if secrets is not None else {"credentials": {"access_token": "github_pat_alice"}}
|
||||
return str(conn.execute(
|
||||
text(
|
||||
"INSERT INTO connector_sessions (user_id, provider, connector_key, auth_kind, status, account_label, "
|
||||
"encrypted_credentials) VALUES (:u, 'github', 'github', :a, :s, :l, :e) RETURNING id"
|
||||
),
|
||||
{"u": user, "a": auth_kind, "s": status, "l": label, "e": encrypt_json(secrets, user)},
|
||||
).scalar())
|
||||
|
||||
|
||||
def _response(payload, status=200, headers=None):
|
||||
response = MagicMock(status_code=status, headers=headers or {})
|
||||
response.json.return_value = payload
|
||||
response.ok = status < 400
|
||||
return response
|
||||
|
||||
|
||||
class TestCatalog:
|
||||
def test_github_syncs_and_reads_with_a_token_and_no_admin_setup(self, monkeypatch):
|
||||
from docsgpt.core.settings import settings
|
||||
|
||||
for name in ("GITHUB_CLIENT_ID", "GITHUB_CLIENT_SECRET", "GITHUB_APP_SLUG"):
|
||||
monkeypatch.setattr(settings, name, None)
|
||||
definition = catalog.get_definition("github")
|
||||
assert definition.configured
|
||||
assert definition.capabilities == ("sync", "read")
|
||||
assert definition.sync_ingestor == "github"
|
||||
assert definition.mcp_url == READONLY_MCP
|
||||
assert definition.setup == {"tools": "ask", "sync": "ask"}
|
||||
payload = definition.to_dict()
|
||||
assert payload["sign_in_methods"] == ["api_key"]
|
||||
assert [f["key"] for f in payload["credential_fields"]] == ["access_token"]
|
||||
|
||||
def test_github_app_sign_in_is_offered_once_configured(self, app_settings):
|
||||
assert catalog.get_definition("github").to_dict()["sign_in_methods"] == ["oauth", "api_key"]
|
||||
|
||||
def test_other_connectors_have_one_sign_in(self):
|
||||
assert catalog.get_definition("google_drive").to_dict()["sign_in_methods"] == ["oauth"]
|
||||
assert catalog.get_definition("s3").to_dict()["sign_in_methods"] == ["api_key"]
|
||||
|
||||
def test_generic_mcp_tools_do_not_belong_to_github(self):
|
||||
"""Every MCP tool would otherwise be listed as GitHub's."""
|
||||
assert catalog.definition_for_tool("mcp_tool") is None
|
||||
assert "mcp_tool" not in catalog.tool_connector_keys()
|
||||
|
||||
|
||||
class TestTokenSignIn:
|
||||
def test_token_is_checked_and_named_after_the_account(self, app, pg_conn):
|
||||
from docsgpt.api.connector.connections import ConnectionsList
|
||||
|
||||
with _db(pg_conn), patch("docsgpt.connectors.github.requests.get",
|
||||
return_value=_response({"login": "octocat"})) as get:
|
||||
resp = _call(app, ConnectionsList, "post", "/api/connections",
|
||||
body={"connector_key": "github", "credentials": {"access_token": "github_pat_abc"}})
|
||||
assert resp.status_code == 201
|
||||
assert resp.get_json()["connection"]["account_label"] == "octocat"
|
||||
assert get.call_args.kwargs["headers"]["Authorization"] == "Bearer github_pat_abc"
|
||||
assert resp.get_json()["setup"] == {"tools": "ask", "sync": "ask"}
|
||||
|
||||
def test_rejected_token_is_not_stored(self, app, pg_conn):
|
||||
from docsgpt.api.connector.connections import ConnectionsList
|
||||
|
||||
with _db(pg_conn), patch("docsgpt.connectors.github.requests.get",
|
||||
return_value=_response({"message": "Bad credentials"}, 401)):
|
||||
resp = _call(app, ConnectionsList, "post", "/api/connections",
|
||||
body={"connector_key": "github", "credentials": {"access_token": "nope"}})
|
||||
assert resp.status_code == 400
|
||||
assert resp.get_json()["code"] == "invalid_credentials"
|
||||
assert pg_conn.execute(text("SELECT count(*) FROM connector_sessions")).scalar() == 0
|
||||
|
||||
|
||||
class TestAppSignIn:
|
||||
def test_callback_stores_the_app_token_under_the_login(self, app, pg_conn, app_settings):
|
||||
import base64
|
||||
import json
|
||||
|
||||
from docsgpt.api.connector.routes import ConnectorsCallback, build_authorization
|
||||
|
||||
with _db(pg_conn), patch("docsgpt.api.connector.routes.service.ensure_can_store_credentials"):
|
||||
started = build_authorization("github", "alice")
|
||||
assert started["authorization_url"].startswith("https://github.com/login/oauth/authorize?")
|
||||
state = started["state"]
|
||||
token_info = {
|
||||
"access_token": "ghu_a", "refresh_token": "ghr_r", "expiry": "2099-01-01T00:00:00+00:00",
|
||||
"user_info": {"login": "octocat", "name": "Octo"},
|
||||
}
|
||||
with _db(pg_conn), patch("docsgpt.parser.connectors.github.auth.GitHubAuth.exchange_code_for_tokens",
|
||||
return_value=token_info):
|
||||
with app.test_request_context(f"/api/connectors/callback?code=c&state={state}"):
|
||||
page = ConnectorsCallback().get()
|
||||
assert page.status_code == 200
|
||||
assert b"github_auth_success" in page.data
|
||||
row = pg_conn.execute(text("SELECT * FROM connector_sessions WHERE provider = 'github'")).one()._mapping
|
||||
assert row["account_label"] == "octocat" and row["auth_kind"] == "oauth"
|
||||
assert service.read_secrets(dict(row))["token_info"]["refresh_token"] == "ghr_r"
|
||||
assert json.loads(base64.urlsafe_b64decode(state))["provider"] == "github"
|
||||
|
||||
def test_installation_link_carries_the_same_state(self, pg_conn, app_settings):
|
||||
from docsgpt.api.connector.routes import build_authorization
|
||||
|
||||
cid = _connection(pg_conn, auth_kind="oauth", secrets={"token_info": {"access_token": "ghu"}})
|
||||
with _db(pg_conn), patch("docsgpt.api.connector.routes.service.ensure_can_store_credentials"):
|
||||
started = build_authorization("github", "alice", cid, install=True)
|
||||
assert started["authorization_url"].startswith(
|
||||
"https://github.com/apps/docsgpt-acme/installations/new?state="
|
||||
)
|
||||
|
||||
def test_installation_redirect_without_state_is_a_friendly_page(self, app, app_settings):
|
||||
"""Installing the app from GitHub itself lands on the callback with no state."""
|
||||
from docsgpt.api.connector.routes import ConnectorsCallback
|
||||
|
||||
with app.test_request_context("/api/connectors/callback?code=c&installation_id=9&setup_action=install"):
|
||||
page = ConnectorsCallback().get()
|
||||
assert page.status_code == 200
|
||||
assert b"installed" in page.data.lower()
|
||||
|
||||
def test_expired_app_token_is_refreshed_before_use(self, pg_conn, app_settings):
|
||||
cid = _connection(pg_conn, auth_kind="oauth", secrets={"token_info": {
|
||||
"access_token": "ghu_old", "refresh_token": "ghr_r", "expiry": "2000-01-01T00:00:00+00:00",
|
||||
}})
|
||||
row = pg_conn.execute(text("SELECT * FROM connector_sessions WHERE id = CAST(:c AS uuid)"),
|
||||
{"c": cid}).one()._mapping
|
||||
with _db(pg_conn), patch(
|
||||
"docsgpt.parser.connectors.github.auth.GitHubAuth.refresh_access_token",
|
||||
return_value={"access_token": "ghu_new", "refresh_token": "ghr_next", "expiry": "2099-01-01T00:00:00+00:00"},
|
||||
):
|
||||
assert service.access_credentials(dict(row)) == {"access_token": "ghu_new"}
|
||||
|
||||
|
||||
class TestRepositories:
|
||||
def test_token_lists_the_repositories_it_can_read(self, app, pg_conn):
|
||||
from docsgpt.api.connector.connections import ConnectionRepositories
|
||||
|
||||
cid = _connection(pg_conn)
|
||||
page = [
|
||||
{"full_name": "octocat/private", "name": "private", "owner": {"login": "octocat"}, "private": True,
|
||||
"description": "Secret", "default_branch": "main", "pushed_at": "2026-09-01T00:00:00Z",
|
||||
"html_url": "https://github.com/octocat/private"},
|
||||
]
|
||||
with _db(pg_conn), patch("docsgpt.connectors.github.requests.get", return_value=_response(page)) as get:
|
||||
resp = _call(app, ConnectionRepositories, "get", f"/api/connections/{cid}/repositories", args=[cid])
|
||||
assert resp.status_code == 200
|
||||
data = resp.get_json()
|
||||
assert data["repositories"] == [{
|
||||
"full_name": "octocat/private", "private": True, "description": "Secret",
|
||||
"default_branch": "main", "updated_at": "2026-09-01T00:00:00Z",
|
||||
"html_url": "https://github.com/octocat/private",
|
||||
}]
|
||||
assert data["install_url"] is None
|
||||
assert get.call_args.args[0] == "https://api.github.com/user/repos"
|
||||
assert "github_pat" not in resp.get_data(as_text=True)
|
||||
|
||||
def test_app_sign_in_lists_the_installations_repositories(self, app, pg_conn, app_settings):
|
||||
from docsgpt.api.connector.connections import ConnectionRepositories
|
||||
|
||||
cid = _connection(pg_conn, auth_kind="oauth", secrets={"token_info": {
|
||||
"access_token": "ghu_a", "expiry": "2099-01-01T00:00:00+00:00"}})
|
||||
|
||||
def fake_get(url, **kwargs):
|
||||
if url.endswith("/user/installations"):
|
||||
return _response({"installations": [{"id": 7}]})
|
||||
assert url.endswith("/user/installations/7/repositories")
|
||||
return _response({"repositories": [{"full_name": "acme/api", "private": True}]})
|
||||
|
||||
with _db(pg_conn), patch("docsgpt.connectors.github.requests.get", side_effect=fake_get):
|
||||
resp = _call(app, ConnectionRepositories, "get", f"/api/connections/{cid}/repositories", args=[cid])
|
||||
data = resp.get_json()
|
||||
assert [r["full_name"] for r in data["repositories"]] == ["acme/api"]
|
||||
assert data["install_url"] == "https://github.com/apps/docsgpt-acme/installations/new"
|
||||
|
||||
def test_rejected_token_asks_to_reconnect(self, app, pg_conn):
|
||||
from docsgpt.api.connector.connections import ConnectionRepositories
|
||||
|
||||
cid = _connection(pg_conn)
|
||||
with _db(pg_conn), patch("docsgpt.connectors.github.requests.get",
|
||||
return_value=_response({"message": "Bad credentials"}, 401)), \
|
||||
patch.object(service, "mark_reconnect_needed") as flag:
|
||||
resp = _call(app, ConnectionRepositories, "get", f"/api/connections/{cid}/repositories", args=[cid])
|
||||
assert resp.status_code == 409
|
||||
flag.assert_called_once()
|
||||
|
||||
def test_only_the_owner_and_only_github(self, app, pg_conn):
|
||||
from docsgpt.api.connector.connections import ConnectionRepositories
|
||||
|
||||
cid = _connection(pg_conn, user="bob")
|
||||
with _db(pg_conn):
|
||||
resp = _call(app, ConnectionRepositories, "get", f"/api/connections/{cid}/repositories", args=[cid])
|
||||
assert resp.status_code == 404
|
||||
|
||||
|
||||
class TestSetup:
|
||||
def test_sync_queues_the_repository_with_the_connection(self, app, pg_conn):
|
||||
from docsgpt.api.connector.connections import ConnectionSetup
|
||||
|
||||
cid = _connection(pg_conn)
|
||||
with _db(pg_conn), patch("docsgpt.api.user.tasks.ingest_remote.apply_async",
|
||||
return_value=MagicMock(id="t")) as apply:
|
||||
resp = _call(app, ConnectionSetup, "post", f"/api/connections/{cid}/setup", body={
|
||||
"create_tools": False,
|
||||
"sync": {"items": {"repo_url": "https://github.com/octocat/private.git"}, "frequency": "daily"},
|
||||
}, args=[cid])
|
||||
assert resp.status_code == 200
|
||||
kwargs = apply.call_args.kwargs["kwargs"]
|
||||
assert kwargs["loader"] == "github"
|
||||
assert kwargs["source_data"] == {"repo_url": "octocat/private"}
|
||||
assert kwargs["connection_id"] == cid
|
||||
assert kwargs["sync_frequency"] == "daily"
|
||||
# Named after the repository, not the connector.
|
||||
assert resp.get_json()["sources"][0]["name"] == "octocat/private"
|
||||
|
||||
def test_not_a_repository_is_a_bad_request(self, app, pg_conn):
|
||||
from docsgpt.api.connector.connections import ConnectionSetup
|
||||
|
||||
cid = _connection(pg_conn)
|
||||
with _db(pg_conn):
|
||||
resp = _call(app, ConnectionSetup, "post", f"/api/connections/{cid}/setup", body={
|
||||
"create_tools": False, "sync": {"items": {"repo_url": "https://example.com/x"}},
|
||||
}, args=[cid])
|
||||
assert resp.status_code == 400
|
||||
|
||||
def test_tools_are_the_read_only_mcp_server_with_discovered_actions(self, app, pg_conn):
|
||||
from docsgpt.api.connector.connections import ConnectionSetup
|
||||
|
||||
cid = _connection(pg_conn)
|
||||
actions = [{"name": "get_file_contents", "description": "Read a file", "annotations": {"readOnlyHint": True},
|
||||
"parameters": {"type": "object", "properties": {"path": {"type": "string"}}}}]
|
||||
with _db(pg_conn), patch("docsgpt.agents.tools.mcp_tool.MCPTool.discover_tools"), \
|
||||
patch("docsgpt.agents.tools.mcp_tool.MCPTool.get_actions_metadata", return_value=actions), \
|
||||
patch("docsgpt.agents.tools.mcp_tool.MCPTool.__init__", return_value=None) as init:
|
||||
for _ in range(2):
|
||||
resp = _call(app, ConnectionSetup, "post", f"/api/connections/{cid}/setup",
|
||||
body={"create_tools": True}, args=[cid])
|
||||
assert resp.status_code == 200
|
||||
# The first MCPTool is the discovery one (the tool registry builds more).
|
||||
config = init.call_args_list[0].args[0]
|
||||
assert config["server_url"] == READONLY_MCP
|
||||
assert config["auth_credentials"] == {"access_token": "github_pat_alice"}
|
||||
rows = pg_conn.execute(text(
|
||||
"SELECT name, config, actions FROM user_tools WHERE connection_id = CAST(:c AS uuid)"
|
||||
), {"c": cid}).all()
|
||||
assert len(rows) == 1
|
||||
name, stored, stored_actions = rows[0]
|
||||
assert name == "mcp_tool"
|
||||
assert stored == {"server_url": READONLY_MCP, "auth_type": "bearer", "transport_type": "http", "timeout": 30}
|
||||
assert "github_pat" not in str(stored)
|
||||
assert stored_actions[0]["name"] == "get_file_contents"
|
||||
assert stored_actions[0]["access"] == "read" and not stored_actions[0].get("require_approval")
|
||||
|
||||
def test_unreachable_mcp_server_creates_nothing(self, app, pg_conn):
|
||||
from docsgpt.api.connector.connections import ConnectionSetup
|
||||
|
||||
cid = _connection(pg_conn)
|
||||
with _db(pg_conn), patch("docsgpt.connectors.mcp._discover", side_effect=RuntimeError("down")):
|
||||
resp = _call(app, ConnectionSetup, "post", f"/api/connections/{cid}/setup",
|
||||
body={"create_tools": True}, args=[cid])
|
||||
assert resp.status_code == 502
|
||||
assert resp.get_json()["code"] == "tools_unavailable"
|
||||
assert pg_conn.execute(text("SELECT count(*) FROM user_tools")).scalar() == 0
|
||||
|
||||
|
||||
class TestToolRuntime:
|
||||
def _tool(self, cid, server_url=READONLY_MCP):
|
||||
return {
|
||||
"id": "tool-gh", "user_id": "alice", "name": "mcp_tool", "connection_id": cid,
|
||||
"config": {"server_url": server_url, "auth_type": "bearer"},
|
||||
"actions": [{"name": "get_me", "active": True}], "credential_mode": "owner",
|
||||
}
|
||||
|
||||
def test_app_token_reaches_the_tool_as_a_bearer_token(self, pg_conn, app_settings):
|
||||
from docsgpt.agents.tool_executor import ToolExecutor
|
||||
|
||||
cid = _connection(pg_conn, auth_kind="oauth", secrets={"token_info": {
|
||||
"access_token": "ghu_live", "expiry": "2099-01-01T00:00:00+00:00"}})
|
||||
with _db(pg_conn), patch("docsgpt.agents.tool_executor.ToolManager") as manager:
|
||||
ToolExecutor(user="alice")._get_or_load_tool(self._tool(cid), "t1", "get_me")
|
||||
config = manager.return_value.load_tool.call_args.kwargs["tool_config"]
|
||||
assert config["auth_credentials"] == {"access_token": "ghu_live"}
|
||||
assert config["server_url"] == READONLY_MCP
|
||||
|
||||
def test_token_never_goes_to_another_server(self, pg_conn):
|
||||
from docsgpt.agents.tool_executor import ToolExecutor
|
||||
|
||||
cid = _connection(pg_conn)
|
||||
with _db(pg_conn), patch("docsgpt.agents.tool_executor.ToolManager") as manager:
|
||||
with pytest.raises(service.ConnectionUnavailable):
|
||||
ToolExecutor(user="alice")._get_or_load_tool(
|
||||
self._tool(cid, "https://evil.example.com/mcp"), "t1", "get_me",
|
||||
)
|
||||
manager.return_value.load_tool.assert_not_called()
|
||||
@@ -0,0 +1,128 @@
|
||||
"""Tests for the GitHub App user sign-in."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import datetime
|
||||
from unittest.mock import MagicMock, patch
|
||||
from urllib.parse import parse_qs, urlparse
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def configured(monkeypatch):
|
||||
from docsgpt.core.settings import settings
|
||||
|
||||
monkeypatch.setattr(settings, "GITHUB_CLIENT_ID", "Iv1.client")
|
||||
monkeypatch.setattr(settings, "GITHUB_CLIENT_SECRET", "app-secret")
|
||||
monkeypatch.setattr(settings, "GITHUB_APP_SLUG", "docsgpt-acme")
|
||||
monkeypatch.setattr(settings, "CONNECTOR_REDIRECT_BASE_URI", "https://docs.example/api/connectors/callback")
|
||||
|
||||
|
||||
def _response(payload, status=200):
|
||||
response = MagicMock(status_code=status)
|
||||
response.json.return_value = payload
|
||||
response.raise_for_status.return_value = None
|
||||
return response
|
||||
|
||||
|
||||
def test_needs_the_app_settings(monkeypatch):
|
||||
from docsgpt.core.settings import settings
|
||||
from docsgpt.parser.connectors.github.auth import GitHubAuth
|
||||
|
||||
monkeypatch.setattr(settings, "GITHUB_CLIENT_ID", None)
|
||||
with pytest.raises(ValueError, match="GITHUB_CLIENT_ID"):
|
||||
GitHubAuth()
|
||||
|
||||
|
||||
def test_authorization_url_carries_state_and_callback(configured):
|
||||
from docsgpt.parser.connectors.github.auth import GitHubAuth
|
||||
|
||||
url = urlparse(GitHubAuth().get_authorization_url(state="st"))
|
||||
query = parse_qs(url.query)
|
||||
assert f"{url.scheme}://{url.netloc}{url.path}" == "https://github.com/login/oauth/authorize"
|
||||
assert query["client_id"] == ["Iv1.client"]
|
||||
assert query["state"] == ["st"]
|
||||
assert query["redirect_uri"] == ["https://docs.example/api/connectors/callback"]
|
||||
|
||||
|
||||
def test_installation_url_carries_state(configured):
|
||||
from docsgpt.parser.connectors.github.auth import GitHubAuth
|
||||
|
||||
assert GitHubAuth().get_installation_url(state="st") == (
|
||||
"https://github.com/apps/docsgpt-acme/installations/new?state=st"
|
||||
)
|
||||
|
||||
|
||||
def test_exchange_returns_expiring_tokens_and_the_login(configured):
|
||||
from docsgpt.parser.connectors.github.auth import GitHubAuth
|
||||
|
||||
token = _response({
|
||||
"access_token": "ghu_a", "refresh_token": "ghr_r", "expires_in": 28800, "token_type": "bearer",
|
||||
})
|
||||
user = _response({"login": "octocat", "name": "The Octocat"})
|
||||
with patch("docsgpt.parser.connectors.github.auth.requests.post", return_value=token) as post, \
|
||||
patch("docsgpt.parser.connectors.github.auth.requests.get", return_value=user):
|
||||
info = GitHubAuth().exchange_code_for_tokens("code-1")
|
||||
assert post.call_args.kwargs["data"]["code"] == "code-1"
|
||||
assert post.call_args.kwargs["headers"]["Accept"] == "application/json"
|
||||
assert info["access_token"] == "ghu_a" and info["refresh_token"] == "ghr_r"
|
||||
assert info["user_info"]["login"] == "octocat"
|
||||
expiry = datetime.datetime.fromisoformat(info["expiry"])
|
||||
assert expiry > datetime.datetime.now(datetime.timezone.utc) + datetime.timedelta(hours=7)
|
||||
|
||||
|
||||
def test_exchange_error_in_a_200_body_raises(configured):
|
||||
"""GitHub reports a bad code with 200 and an ``error`` field."""
|
||||
from docsgpt.parser.connectors.github.auth import GitHubAuth
|
||||
|
||||
with patch("docsgpt.parser.connectors.github.auth.requests.post",
|
||||
return_value=_response({"error": "bad_verification_code"})):
|
||||
with pytest.raises(ValueError, match="bad_verification_code"):
|
||||
GitHubAuth().exchange_code_for_tokens("stale")
|
||||
|
||||
|
||||
def test_refused_refresh_raises(configured):
|
||||
from docsgpt.parser.connectors.github.auth import GitHubAuth
|
||||
|
||||
with patch("docsgpt.parser.connectors.github.auth.requests.post",
|
||||
return_value=_response({"error": "bad_refresh_token"})):
|
||||
with pytest.raises(ValueError, match="bad_refresh_token"):
|
||||
GitHubAuth().refresh_access_token("ghr_old")
|
||||
|
||||
|
||||
def test_refresh_rotates_the_refresh_token(configured):
|
||||
from docsgpt.parser.connectors.github.auth import GitHubAuth
|
||||
|
||||
with patch("docsgpt.parser.connectors.github.auth.requests.post", return_value=_response({
|
||||
"access_token": "ghu_b", "refresh_token": "ghr_new", "expires_in": 28800,
|
||||
})) as post:
|
||||
info = GitHubAuth().refresh_access_token("ghr_old")
|
||||
assert post.call_args.kwargs["data"]["grant_type"] == "refresh_token"
|
||||
assert info["access_token"] == "ghu_b" and info["refresh_token"] == "ghr_new"
|
||||
|
||||
|
||||
class TestExpiry:
|
||||
def test_token_without_expiry_never_expires(self, configured):
|
||||
"""An app with token expiry turned off issues tokens with no expiry: they stay valid."""
|
||||
from docsgpt.parser.connectors.github.auth import GitHubAuth
|
||||
|
||||
assert GitHubAuth().is_token_expired({"access_token": "ghu_a"}) is False
|
||||
|
||||
def test_expired_and_fresh(self, configured):
|
||||
from docsgpt.parser.connectors.github.auth import GitHubAuth
|
||||
|
||||
now = datetime.datetime.now(datetime.timezone.utc)
|
||||
auth = GitHubAuth()
|
||||
assert auth.is_token_expired({"access_token": "a", "expiry": (now - datetime.timedelta(minutes=1)).isoformat()})
|
||||
assert not auth.is_token_expired({"access_token": "a", "expiry": (now + datetime.timedelta(hours=1)).isoformat()})
|
||||
assert auth.is_token_expired({})
|
||||
|
||||
|
||||
def test_registered_as_an_auth_provider_but_not_a_file_connector():
|
||||
"""GitHub signs in like the OAuth connectors but ingests through the remote loader."""
|
||||
from docsgpt.parser.connectors.connector_creator import ConnectorCreator
|
||||
|
||||
assert ConnectorCreator.has_auth("github")
|
||||
assert not ConnectorCreator.is_supported("github")
|
||||
assert "github" not in ConnectorCreator.get_supported_connectors()
|
||||
Reference in new issue
Block a user