mirror of
https://github.com/tiennm99/DocsGPT.git
synced 2026-10-11 03:12:55 +00:00
Merge pull request #2843 from arc53-machine/connectors
Connectors: one place to connect accounts for Sources and Tools
This commit is contained in:
317 files changed
+44582
-3590
No files matched your search
@@ -68,6 +68,8 @@ A multi-phase agent designed for in-depth research tasks:
|
||||
|
||||
Includes budget controls for max steps, timeout, and token limits to keep research bounded.
|
||||
|
||||
Research steps can't stop to ask you anything, so a tool action that needs approval, a connection or the user's app isn't run there; the step is told why and carries on. The same goes for actions an API, widget or public-link user can't take on the owner's accounts.
|
||||
|
||||
**Best for:** Complex questions that require multi-step investigation, gathering information from multiple sources, and producing structured reports with citations.
|
||||
|
||||
### 4. Workflow Agent
|
||||
@@ -111,6 +113,7 @@ Once an agent is created, you can:
|
||||
* Modify any of its configuration settings (name, description, source, prompt, tools, type).
|
||||
* **Generate a Public Link:** From the edit screen, you can create a shareable public link that allows others to import and use your agent.
|
||||
* **Get a Webhook URL:** You can also obtain a Webhook URL for the agent. This allows external applications or services to trigger the agent and receive responses programmatically, enabling powerful integrations and automations.
|
||||
* **Share it with a team:** Viewers chat with it; editors can also change it. The share dialog lists what the agent uses and whose access each item runs with. See [Sharing agents and what they use](/Deploying/Access-Control#sharing-agents-and-what-they-use).
|
||||
* **Use it via API:** Every agent exposes an API key that can be used with the native [Agent API](/Agents/api) or the [OpenAI-Compatible API](/Agents/openai-compatible) so you can drop DocsGPT Agents into any tool that already speaks the chat completions protocol.
|
||||
|
||||
## Seeding Premade Agents from YAML
|
||||
|
||||
@@ -172,7 +172,7 @@ Open the agent in the builder and expand the **Guardrails** section.
|
||||
Checks that cannot run without settings (`denylist`, `url`, `policy`, and `pii` with no entities selected) are marked *Not configured* and block saving until they are filled in, so a half-configured control can never be published as if it were protecting you.
|
||||
|
||||
<Callout type="info" emoji="ℹ️">
|
||||
Guardrails are the **agent owner's** policy. Team members with edit access can see the configuration but cannot change it — an editor who could clear a control would silently strip protection from everyone using the agent. Controls required by the instance floor (below) appear locked and cannot be removed.
|
||||
Guardrails, like the agent's limits, are policy that the owner and the agent's team **editors** can change. Viewers don't see the configuration. To keep a control that nobody who edits the agent can remove, require it in the [instance floor](#the-instance-floor): floor controls appear locked for everyone. See [Sharing agents and what they use](/Deploying/Access-Control#sharing-agents-and-what-they-use) for what else viewers and editors can do.
|
||||
</Callout>
|
||||
|
||||
Guardrails also apply to a **draft** agent in the builder preview, which is the natural place to try a control before publishing.
|
||||
|
||||
@@ -138,17 +138,95 @@ Four resource types can be shared: **agents**, **sources**, **prompts**, and **t
|
||||
|
||||
Sharing rules:
|
||||
|
||||
- Only the **owner** of a resource can share it.
|
||||
- Only the **owner** of a resource can share it, unless they turn on **Editors can share** in the share dialog's **Access settings**.
|
||||
- A share targets either the **whole team** or a **single member**.
|
||||
- Each share carries an access level: **`viewer`** (read-only) or **`editor`** (read and modify).
|
||||
- `editor` is not the same as owner — an editor can change a resource but cannot delete it or re-share it.
|
||||
- Shared tools run server-side with the **owner's** credentials; a grantee never sees the owner's secrets.
|
||||
- `editor` is not the same as owner — an editor can change a resource but cannot re-share it unless the owner turns on **Editors can share**, or delete it unless the owner turns on **Editors can delete** (agents and sources only; editors never delete tools or prompts).
|
||||
- Shared tools run server-side with the **owner's** credentials, or with each member's own account when a connected tool is shared that way; a grantee never sees the owner's secrets.
|
||||
- A wiki's editors can edit its pages, but only its owner decides whether API and widget users can edit it through an agent (see [Wiki sources](/Sources/Wiki-sources#edits-from-the-api-widget-and-public-links)).
|
||||
|
||||
### Sharing agents and what they use
|
||||
|
||||
Sharing an agent lets people use it; it doesn't share its tools, sources or prompt. Each stays with whoever owns it, and people reach them only through the agent.
|
||||
|
||||
- **Viewers** chat with the agent. They don't see its configuration, and they never open its share dialog.
|
||||
- **Editors** open its edit page and change it: its instructions, model, tools, sources and prompt, and its [guardrails](/Agents/guardrails) and limits. They can share it or delete it only when you turn on **Editors can share** or **Editors can delete**.
|
||||
|
||||
Every tool, source and prompt on the agent runs with one person's access for everyone who uses the agent. The exception is a connected tool shared as **Each person's own**, which uses the account of whoever is chatting. The agent's share dialog lists them under **What this agent uses**, for you and its editors. For a workflow agent, the list also includes the tools and sources on its nodes.
|
||||
|
||||
| The item | Runs with | The list says |
|
||||
| --- | --- | --- |
|
||||
| Yours, or shared with you | Your access | **Your access** (an editor sees **The owner's access**) |
|
||||
| Added by an editor while you can't use it | The editor's access; they are its [sponsor](#resources-an-editor-adds-to-someone-elses-agent). Once you can use it, it runs with yours again | **dana@example.com's access** |
|
||||
| A tool with saved credentials (an API key, an MCP sign-in) | Its owner's credentials | **Your saved credentials**, or **dana@example.com's saved credentials** for a teammate's tool |
|
||||
| A connected tool shared as **Your account** | That account on the service | **Your Notion account**, or **dana@example.com's Notion account** for a teammate's tool |
|
||||
| A connected tool shared as **Each person's own** | The account of whoever is chatting; they connect it the first time. API and widget users run the agent as its owner, so they get the owner's account | **Each person's own Notion account (API and widget: yours)** |
|
||||
|
||||
The list and the edit page name a person only when you know them already: you, the agent's owner, anyone who shares a team with you and, when you own the agent, whoever sponsored something on it. An item's owner is named for its account or credentials, or as the person to ask about it, only when you can also see that item. Anyone else shows as someone else, for example **Someone else's access**.
|
||||
|
||||
An editor becomes a sponsor only after confirming who will reach the resource through the agent. A resource whose sponsor or owner loses access stops running and is marked **Stopped** in the list, with the reason (see [When a resource stops working](#when-a-resource-stops-working)).
|
||||
|
||||
Some runs are limited, because nobody can approve an action for you there:
|
||||
|
||||
- **API key, website widget and public link.** A write action on your connected accounts or saved credentials runs only if you allow it under **Access details > Actions API, widget and public-link users can take as you**. Only the owner can change that list. It offers every such write on the agent, including tools an editor sponsored and the tools on a workflow's nodes. The share dialog marks a tool **API changes off** when none of its writes are allowed, **Some API changes off** when only some are, and **Changes off by an admin** when an admin turned off changes through its service. On a tool shared as **Each person's own**, public-link users act with their own account, so the list doesn't limit them there; a write that needs approval still asks them first. See [Agents used through an API key](/Guides/Connectors#agents-used-through-an-api-key).
|
||||
- **Wikis.** API and widget users edit a wiki only when its owner turns on **Let API and widget users edit this wiki** in its **Wiki settings**. Public-link users edit only wikis they can edit themselves, and approve each edit (see [Wiki sources](/Sources/Wiki-sources#edits-from-the-api-widget-and-public-links)).
|
||||
- **Research agents.** A research step can't stop to ask, so it skips any action that would need approval or a connection, and any write the caller may not make on your accounts. The step is told why and carries on (see [Research Agent](/Agents/basics#3-research-agent)).
|
||||
|
||||
### Resources an editor adds to someone else's agent
|
||||
|
||||
An agent (and its workflow) runs as its owner, so its tools, sources and prompts are checked against the owner's access. When a team editor adds one the owner can't use, it runs with the **editor's** access instead, for everyone who uses the agent: members of the teams it is shared with, anyone with its API key or website widget, its public link and its webhook. The editor becomes that resource's **sponsor**.
|
||||
|
||||
- Only someone who **owns** the resource, or has **`editor`** access to it through any team, can sponsor it. `viewer` access lets you use a resource in your own agents, but not extend it to another agent's users: the save is refused with `403` and `code: "sponsor_not_allowed"`.
|
||||
- Sponsoring is never implied. A save that would make you a new sponsor is refused with `409` until you confirm it. In DocsGPT, a dialog names each resource and who will reach it through the agent. Through the API, send the save again with `confirm_sponsor` listing every resource from the response as `"<type>:<id>"`. A `confirm_sponsor` entry for anything the save doesn't ask you to sponsor is refused with `400`.
|
||||
- A sponsored resource stops running when its sponsor can no longer edit the agent, or no longer owns or edits the resource. It stays on the agent but does nothing, and it doesn't pass to whoever saves the agent next. The agent's edit page says why it stopped. An editor who may sponsor it can choose **Run … with my access** there; after they confirm in the same dialog that names who reaches the agent, it runs with their access from their next save. Through the API, include its key in `confirm_sponsor` on any save. Otherwise, someone removes it.
|
||||
- A resource that is removed and later added again needs a new confirmation, even if its old sponsor could still sponsor it.
|
||||
- Editors can remove any tool, source or prompt from the agent or its workflow nodes, including the owner's private ones they can't open.
|
||||
- Workflow nodes follow the same rules, through `PUT /api/workflows/<id>`. The owner's own saves are checked too: a node can't name a tool or source its owner can't use.
|
||||
- The confirmation covers the audience the agent has at that moment. If the owner later shares the agent with more teams, or turns on a public link, sponsored resources reach those people too, and their sponsors aren't asked again.
|
||||
- `audience.teams` lists every team the agent is shared with, including teams where only some members were given access.
|
||||
|
||||
A save that needs confirmation returns:
|
||||
|
||||
```json
|
||||
{
|
||||
"success": false,
|
||||
"code": "sponsor_confirmation_required",
|
||||
"message": "These resources would run with your access for everyone who uses this agent. Confirm to add them.",
|
||||
"resources": [{ "key": "tool:<id>", "type": "tool", "id": "<id>", "name": "Jira" }],
|
||||
"audience": { "teams": ["Support"], "api_key": true, "public_link": false, "webhook": false }
|
||||
}
|
||||
```
|
||||
|
||||
Owners and editors see sponsored resources on the agent's edit page and in `resource_sponsors` from `GET /api/get_agent` (and `GET /api/workflows/<id>`). Viewers get an empty list. Each entry has the resource (`key`, `type`, `id`, `name`), the sponsor (`user_id`, `label`, both `null` for someone you don't know; see [Sharing agents and what they use](#sharing-agents-and-what-they-use)), `state` (`active` or `inactive`), `reason` when inactive (`sponsor_cannot_edit_agent` or `sponsor_cannot_edit_resource`), and `can_confirm`, which says whether you may take an inactive one over.
|
||||
|
||||
<Callout type="warning" emoji="⚠️">
|
||||
After upgrading, resources sponsored by someone with only `viewer` access to them stop running. An editor who owns or can edit such a resource can choose **Run … with my access** on the agent's edit page to start it again.
|
||||
</Callout>
|
||||
|
||||
### When a resource stops working
|
||||
|
||||
Every run checks each tool, source and prompt on the agent (and each tool and source on its workflow nodes) against the owner's access, or the sponsor's. One that no longer passes is left out of the run; a prompt falls back to the default prompt. A connected tool whose account needs attention is different: it stays on the agent but can't run until someone fixes the account. In a chat, calling it shows a card asking to connect the account; in a scheduled run, a webhook or an API call, the call is refused. The agent's edit page, and the workflow builder for node resources, lists each resource that stopped or can't run, with the reason and what you can do:
|
||||
|
||||
| Reason | What happened | What you can do |
|
||||
| --- | --- | --- |
|
||||
| `deleted` | The tool was deleted. Left out of runs. | Remove it. |
|
||||
| `owner_lost_access` | The owner can no longer use it, for example a team stopped sharing it with them. Left out of runs. | Ask its owner to share it again, choose **Run … with my access** if you may sponsor it, or remove it. |
|
||||
| `sponsor_cannot_edit_agent`, `sponsor_cannot_edit_resource` | Its sponsor lost access (see above). Left out of runs. | Choose **Run … with my access** if you may sponsor it, or remove it. |
|
||||
| `connection_needs_reconnect` | The account the tool uses was disconnected or needs signing in again. It can't run until someone signs in again. | If it is your account, choose **Reconnect**; otherwise ask the tool's owner. |
|
||||
| `connection_removed` | The tool's connection was removed but the tool was kept, and it has no credentials of its own. It can't run until the account is connected again. | Ask the tool's owner to connect the account again, or remove it. |
|
||||
| `connector_disabled` | An admin turned the service off. It can't run until it is turned back on. | Ask an admin to turn it back on, or remove it. |
|
||||
|
||||
Only tools that run on their owner's account (owner mode) are checked this way. A tool shared in member mode runs on each person's own account, so the owner's account doesn't decide whether it runs: it is listed as running, with the note that each person uses their own account, and whoever runs it is asked to connect their own account when they need to.
|
||||
|
||||
**Remove** takes the item off the form; save to store it. Through the API, `GET /api/get_agent` and `GET /api/workflows/<id>` return `resource_states` to owners and editors (an empty list to everyone else): one entry per attached resource with `key`, `type`, `id`, `name`, `state` (`active` or `stopped`), `reason`, `note` (`per_user_account` for a running member-mode tool), `sponsor` (`{user_id, label}`), `contact_role` and `contact`, `connection`, `can_confirm` and `can_reconnect`. `contact_role` is `resource_owner` when the resource's owner can fix it; `contact` names them (`{user_id, label}`) when you know them, and is `null` otherwise. `connection` names the service; its `id` comes only with `can_reconnect`, and the account's own name only for its owner. `runs_as` is the live sponsor a running resource runs as (`null` when it runs as the owner). A running tool also has `credential_mode` (`owner` or `member` for a connected tool, else `null`), `account` (whose saved credentials or `owner`-mode connection it uses), `owner_credential_writes` (its write actions on those credentials, which the API write allowlist covers) and `writes_allowed` (`false` when an admin turned off changes through its service). When anything can be taken over, the response also carries `sponsor_audience`, the same shape as `audience` above. People are named by the rule in [Sharing agents and what they use](#sharing-agents-and-what-they-use): `sponsor`, `runs_as` and `account` have `user_id` and `label` set to `null` for someone you don't know, and `contact` is `null`. A resource's name is shown only for a resource that runs, that someone sponsored, that you can see yourself, or that is on an agent (agent saves check every reference). If the run state can't be worked out, the read still succeeds with an empty list.
|
||||
|
||||
The state comes from the checks a run uses, so a source, prompt or tool the page shows as running is one a run uses, and one shown as stopped is left out of runs or can't run until it's fixed, as the table says. Each resource a run leaves out is logged as `resource_stopped` with the agent or workflow, the resource's type and id, and the reason.
|
||||
|
||||
## Audit log
|
||||
|
||||
Access-control actions are appended to the `auth_events` table alongside the [authentication events](/Deploying/OIDC-SSO#login-auditing). This includes admin actions — `admin_user_activated` / `admin_user_deactivated`, `admin_sessions_revoked`, `role_granted` / `role_revoked` (with `metadata.source` = `manual` or `oidc_group`), `quota_policy_set` / `quota_policy_deleted` — and team events (`team.create`, `team.member_add`, `team.member_role`, `team.member_remove`, `team.share`, `team.unshare`, `team.transfer_owner`, `team.delete`).
|
||||
|
||||
Data-plane actions are recorded too: `source.created` / `source.deleted` / `source.reingested`, `agent.created` / `agent.updated` / `agent.deleted` / `agent.key_regenerated`, and `conversation.deleted` / `conversation.deleted_all`.
|
||||
Data-plane actions are recorded too: `source.created` / `source.deleted` / `source.reingested` / `source.wiki_settings_updated`, `agent.created` / `agent.updated` / `agent.deleted` / `agent.key_regenerated`, and `conversation.deleted` / `conversation.deleted_all`.
|
||||
|
||||
Every row carries `actor_id` (who did it) and `target_id` (the user it was done to, or `NULL` when the event is not about a user), so "everything this admin did" is a single query:
|
||||
|
||||
|
||||
@@ -34,7 +34,13 @@ Signing key for session tokens and other signed capabilities. Required on every
|
||||
|
||||
Type `str`, default `default-docsgpt-encryption-key`.
|
||||
|
||||
Key used to encrypt stored credentials such as tool and connector secrets.
|
||||
Key used to encrypt stored credentials such as tool and connector secrets. Set your own value before connecting services on a multi-user install; the default is public.
|
||||
|
||||
### `ENCRYPTION_SECRET_KEY_PREVIOUS`
|
||||
|
||||
Type `str`, default unset.
|
||||
|
||||
Previous ENCRYPTION_SECRET_KEY, tried when a stored credential was encrypted with it. Set it while rotating the key, run `docsgpt connectors reencrypt`, then remove it.
|
||||
|
||||
### `INTERNAL_KEY`
|
||||
|
||||
@@ -1134,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`
|
||||
|
||||
|
||||
@@ -0,0 +1,264 @@
|
||||
---
|
||||
title: Connectors
|
||||
description: Connect DocsGPT to the services your team uses. One connection can sync content into Knowledge and give agents tools, with credentials encrypted on the server.
|
||||
---
|
||||
|
||||
import { Callout } from 'nextra/components'
|
||||
import { Steps } from 'nextra/components'
|
||||
|
||||
# Connectors
|
||||
|
||||
A connector is a service DocsGPT can connect to: Google Drive, SharePoint, Confluence, GitHub, Amazon S3, Reddit, Brave Search, Telegram, ntfy, PostgreSQL, curated MCP servers (Notion, Linear, Atlassian, Sentry, Asana, Stripe), and any MCP server or OpenAPI spec you add yourself. A **connection** is one signed-in account or one saved API key for a connector.
|
||||
|
||||
Connections are the single place credentials live. A synced source and an agent tool both point at a connection, so you sign in once, reconnect once when a token expires, and disconnect in one place.
|
||||
|
||||
## Using connectors
|
||||
|
||||
Open **Settings > Connectors**. Filters narrow the list by category, and **Connected** shows what you already use.
|
||||
|
||||
1. Pick a connector. Services with OAuth open the provider's sign-in page in a pop-up; API-key services ask for the key.
|
||||
2. Choose what the connection sets up. Tool services create their tools right away. Services that can sync ask whether to **Sync into Knowledge**: off by default, on when you started from **Add knowledge** or from Knowledge's **Connect a service**. Left off, the account is still connected and you can sync later with **Sync more content** on its page. Turned on, pick files or folders, how often to sync (never, daily, weekly or monthly) and, under **Advanced retrieval settings**, the same chunking and retrieval options as an upload.
|
||||
3. Adjust what the tools may do. Each action is **Read** or **Write**, and each can be **Always allow**, **Needs approval** or **Off**. Writes default to needing approval.
|
||||
|
||||
**Knowledge** (formerly Sources) is what the assistant searches: uploads, links, and content synced from your connections. The **Add knowledge** and **Add Tool** dialogs, the composer's Knowledge and Tools menus, and the agent builder all start the same flow, and list tools grouped by the connection they use. A service offered two ways appears as one card: Confluence syncs pages into Knowledge and, through **Jira & Confluence**, lets agents search and update Jira and Confluence. GitHub and Linear do both from one connection: see [GitHub](#github) and [Linear](#linear).
|
||||
|
||||
A connector's page is the one place to manage it: its accounts (**Reconnect**, **Rename**, **Disconnect**), the Knowledge it syncs (**Sync now**), and its tools (on or off, and what each action may do). On the Knowledge and Tools pages, anything from a connection has **Manage connection** in its menu, which opens that page. Disconnecting keeps synced content but stops its sync.
|
||||
|
||||
### Several accounts of one service
|
||||
|
||||
You can connect the same service more than once, for example two Telegram bots. Give each account a name when you connect it (**Name this account**) or later with **Rename**. When you have more than one account of a service, its tools are listed as "Telegram · Alerts bot", and the assistant sees each account's actions under that account's name, so it can pick the right one.
|
||||
|
||||
### Fixed values
|
||||
|
||||
Under **Customize**, each action lists its parameters. A parameter is either left to the AI (**Let AI decide**) or set to **Always use** a value. A fixed value is sent on every call; the AI never sees it and can't change it, even when a prompt asks it to. Only the tool's owner can change fixed values.
|
||||
|
||||
A Telegram connection can hold a **Default chat ID**. When it's set, messages always go to that chat. To find a chat's id, add the bot to the chat, send it a message, and open `https://api.telegram.org/bot<token>/getUpdates`. The same bot with another default chat is a separate connection. When team members use their own accounts, each member's messages go to their own default chat.
|
||||
|
||||
### When a connection needs attention
|
||||
|
||||
If a provider rejects a token that DocsGPT cannot refresh, the connection is marked **Reconnect needed**, its synced sources pause, and you get a notification. Reconnecting the same account resumes them. In a chat, a tool whose connection is missing or expired shows a **Connect** prompt instead of failing; the answer continues once you connect.
|
||||
|
||||
### Sharing a tool with a team
|
||||
|
||||
Sharing a tool that uses a connection asks whose account team members use:
|
||||
|
||||
- **Your account.** Everyone acts as your account on that service. When the tool can write, you confirm this before sharing, and a member's write actions always need approval.
|
||||
- **Each person's own.** Each person uses their own account, and sees a **Connect** prompt the first time they use the tool. Through an agent's API key or widget, the agent runs as its owner, so the tool uses the agent owner's account.
|
||||
|
||||
An editor you let share the tool sees your choice as **The owner's account** and can't change it. An agent's share dialog shows the same words for each connected tool on the agent (see [Sharing agents and what they use](/Deploying/Access-Control#sharing-agents-and-what-they-use)).
|
||||
|
||||
Tools from OAuth MCP servers default to each member's own account. An admin can force one mode for a connector.
|
||||
|
||||
### Agents used through an API key
|
||||
|
||||
An agent called with its API key (the website widget, the API) runs with its owner's connections, and nobody can approve an action there. It can read through those connections, but it can't take write actions on them unless the owner allows each one under **Access details > Actions API, widget and public-link users can take as you**. The same goes for tools that hold the owner's own credentials without a connection: an API tool action that sends a header or query value the owner saved (a key or token), or an MCP server the owner signed in to. The list offers every such write on the agent, including tools a team editor sponsored and the tools on a workflow's nodes. Only the owner can change that list; team editors can't. The owner previewing their own agent in DocsGPT is not limited.
|
||||
|
||||
The same list covers people who open the agent from its public link without being on a team it's shared with, and anything they schedule from that chat. They can't approve write actions on the owner's accounts or credentials, so those run only when the owner allows them there. On a tool where each person connects their own account, they act with their own account, so the list doesn't limit them there; a write that needs approval still asks them first.
|
||||
|
||||
Schedules that an API or widget chat set up before this version still run as the owner's own, without this list, until they fire or expire. Schedules set from a public link follow the list from their next run.
|
||||
|
||||
Wikis have their own switch rather than an entry in this list: API and widget users can read a wiki the agent uses but edit it only when the wiki's owner turns on **Let API and widget users edit this wiki** in its **Wiki settings**. Public-link visitors edit only wikis they can edit themselves, and approve each edit (see [Wiki sources](/Sources/Wiki-sources#edits-from-the-api-widget-and-public-links)).
|
||||
|
||||
A webhook is the owner's own automation (its URL is a secret), so its runs are not limited by this list. Nobody can approve during a webhook run, so it still skips every action that needs approval.
|
||||
|
||||
### What answers show
|
||||
|
||||
Tool calls name the service that ran them: "Searched Notion", "Read from Google Drive", "Used Linear: create issue". Citations from a synced source read "From Google Drive". Answers never show which account was used.
|
||||
|
||||
## Admin setup
|
||||
|
||||
Admins manage connectors in **Admin > Connectors**. It lists every connector with its status, how many connections use it, and:
|
||||
|
||||
- **Enabled** turns a connector off for everyone. Members no longer see it (except to manage a connection they already have), its tools stop working, and its sources stop syncing until you turn it back on. Existing connections are kept.
|
||||
- **Shared tools use** is **The sharer decides per share** (default), **Always the sharer's account** or **Always each person's own account**. Connectors that only sync content have no tools and show **No tools**.
|
||||
- The **MCP server** row decides whether members can add their own MCP servers. Presets are switched one by one.
|
||||
- **Setup guide** on an OAuth connector shows the redirect URI to register and which server settings are still missing. GitHub works without setup and shows **Tokens only** until its optional GitHub App settings are in place.
|
||||
- **Write access** has one switch per connector whose tools can opt into changes, currently **Let agents make changes through GitHub** (on by default). Turned off, members no longer see the option, the API refuses it, and every GitHub tool only reads, including ones already set up for changes: they call the read-only endpoint and their write actions are refused with a reason.
|
||||
|
||||
A connector that still needs server settings starts turned off and is hidden from members; its switch stays disabled until the settings are present. It turns on once they are, unless you switched it off. On a phone, the page lists the connectors and opens each one's controls in a panel.
|
||||
|
||||
### Encryption key
|
||||
|
||||
Credentials are encrypted with a key derived from `ENCRYPTION_SECRET_KEY` and bound to the connection's owner. **Set your own value before anyone connects a service.**
|
||||
|
||||
```env
|
||||
ENCRYPTION_SECRET_KEY=a-long-random-value
|
||||
```
|
||||
|
||||
When authentication is on (`AUTH_TYPE` set), DocsGPT refuses to store new credentials while the key is the public default, and Admin > Connectors shows a warning. A single-user local install keeps working and logs a warning at startup.
|
||||
|
||||
To rotate the key, move the old value to `ENCRYPTION_SECRET_KEY_PREVIOUS`, set the new one, restart the API and worker, then run:
|
||||
|
||||
```bash
|
||||
docsgpt connectors reencrypt
|
||||
```
|
||||
|
||||
The command prints how many connections it rewrote. Connections it cannot decrypt with either key are marked **Reconnect needed** for their owners. Once it has run, you can remove `ENCRYPTION_SECRET_KEY_PREVIOUS`.
|
||||
|
||||
### Redirect URIs
|
||||
|
||||
OAuth connectors return to one callback, set by `CONNECTOR_REDIRECT_BASE_URI` (default `http://127.0.0.1:7091/api/connectors/callback`). Register it with each provider exactly as set, without query parameters. MCP servers use `MCP_OAUTH_REDIRECT_URI`, which is derived from the same base when unset. Admin > Connectors shows both values with copy buttons.
|
||||
|
||||
If the frontend runs on a different origin from the API, list it in `CONNECTOR_ALLOWED_ORIGINS` so the sign-in pop-up can hand the result back.
|
||||
|
||||
### Google Drive
|
||||
|
||||
<Steps>
|
||||
|
||||
### Create OAuth credentials
|
||||
|
||||
In the [Google Cloud Console](https://console.cloud.google.com/), enable the **Google Drive API**, then create an **OAuth client ID** of type **Web application**. Add the redirect URI from Admin > Connectors under **Authorized redirect URIs**.
|
||||
|
||||
### Set the server settings
|
||||
|
||||
```env
|
||||
GOOGLE_CLIENT_ID=your-client-id
|
||||
GOOGLE_CLIENT_SECRET=your-client-secret
|
||||
```
|
||||
|
||||
To offer Google's own file picker, also build the frontend with `VITE_GOOGLE_CLIENT_ID` (and optionally `VITE_GOOGLE_PICKER_API_KEY`). Without them, members browse files in DocsGPT's picker.
|
||||
|
||||
### Publish the app
|
||||
|
||||
Publish the OAuth consent screen, or make it an **Internal** Workspace app.
|
||||
|
||||
<Callout type="warning" emoji="⚠️">
|
||||
Apps left in **Testing** get refresh tokens that expire after seven days, which stops background sync.
|
||||
</Callout>
|
||||
|
||||
</Steps>
|
||||
|
||||
### SharePoint and OneDrive
|
||||
|
||||
Register an app in [Microsoft Entra ID](https://entra.microsoft.com/) with a **Web** redirect URI from Admin > Connectors, create a client secret, and grant the delegated Microsoft Graph permissions `Files.Read`, `Sites.Read.All` and `User.Read`.
|
||||
|
||||
```env
|
||||
MICROSOFT_CLIENT_ID=your-application-id
|
||||
MICROSOFT_CLIENT_SECRET=your-client-secret
|
||||
MICROSOFT_TENANT_ID=common # or your tenant id for a single-tenant app
|
||||
```
|
||||
|
||||
See [SharePoint / OneDrive](/Guides/Integrations/sharepoint-connector) for tenant options.
|
||||
|
||||
### Confluence
|
||||
|
||||
Create an **OAuth 2.0 (3LO)** app in the [Atlassian developer console](https://developer.atlassian.com/console/myapps/), add the redirect URI from Admin > Connectors as its callback URL, and add the scopes `read:page:confluence`, `read:space:confluence`, `read:attachment:confluence` and `read:me`. See [Confluence](/Guides/Integrations/confluence-connector) for details.
|
||||
|
||||
```env
|
||||
CONFLUENCE_CLIENT_ID=your-client-id
|
||||
CONFLUENCE_CLIENT_SECRET=your-client-secret
|
||||
```
|
||||
|
||||
### GitHub
|
||||
|
||||
One GitHub connection syncs repositories into Knowledge and gives agents GitHub's own [MCP server](https://github.com/github/github-mcp-server) as a tool. When connecting, members choose both on one screen: **Let agents use GitHub** (on by default) and **Sync into Knowledge** with a repository to sync. A connection set up without tools can add them later from its page.
|
||||
|
||||
- **Knowledge.** The repository picker lists what the connection can read. Each synced repository is one source, read with the connection's token and synced on the chosen schedule.
|
||||
- **Tools.** By default the tool is the read-only endpoint `https://api.githubcopilot.com/mcp/readonly`: agents can read code, issues, pull requests and more, and cannot change anything unless the connection [lets them make changes](#letting-agents-make-changes). Its actions are discovered when the tool is created; **Refresh tools** re-reads them.
|
||||
|
||||
Members connect in one of two ways.
|
||||
|
||||
#### A personal access token (no admin setup)
|
||||
|
||||
Always available. Create a [fine-grained personal access token](https://github.com/settings/personal-access-tokens/new):
|
||||
|
||||
<Steps>
|
||||
|
||||
### Choose the repositories
|
||||
|
||||
Under **Repository access**, pick **Only select repositories** (or all repositories) for the owner whose repositories DocsGPT should read.
|
||||
|
||||
### Grant read access
|
||||
|
||||
Under **Repository permissions**, set **Contents** to **Read-only**. **Metadata** is read-only by default. Add read-only access to **Issues** and **Pull requests** if agents should read those through the tools, or **Read and write** if agents should also [make changes](#letting-agents-make-changes).
|
||||
|
||||
### Paste it into DocsGPT
|
||||
|
||||
DocsGPT checks the token with GitHub and names the connection after the account. A token that expires, or that GitHub stops accepting, marks the connection **Reconnect needed**; reconnect with a new token.
|
||||
|
||||
</Steps>
|
||||
|
||||
A classic token with the `repo` scope also works, but it can read every repository the account can, so prefer a fine-grained one.
|
||||
|
||||
#### Sign in with GitHub (a GitHub App)
|
||||
|
||||
When an admin registers a GitHub App, members can also **Sign in with GitHub**. Each member then chooses the repositories on GitHub, when they install the app, and DocsGPT reads only those. Tokens last eight hours and are renewed automatically.
|
||||
|
||||
<Steps>
|
||||
|
||||
### Register the app
|
||||
|
||||
In GitHub, open **Settings > Developer settings > GitHub Apps > New GitHub App** (under an organization's settings to let its members install it).
|
||||
|
||||
- **Callback URL**: the redirect URI from Admin > Connectors (`CONNECTOR_REDIRECT_BASE_URI`).
|
||||
- Turn on **Request user authorization (OAuth) during installation**. Choosing repositories from DocsGPT then returns to DocsGPT, which reloads the repository list.
|
||||
- Leave **Expire user authorization tokens** on (DocsGPT refreshes them).
|
||||
- **Webhook**: turn off **Active**; DocsGPT does not use webhooks.
|
||||
- **Repository permissions**: **Contents** read-only (**Metadata** read-only is added automatically). Add **Issues** and **Pull requests** read-only for the tools, or read and write if members should be able to let agents [make changes](#letting-agents-make-changes). A member's token can only do what both the app and the member may do, so an app without write permissions keeps every agent read-only. When you add permissions later, each installation has to accept them on GitHub first.
|
||||
- **Where can this GitHub App be installed?**: **Any account** for members outside your organization, otherwise **Only on this account**.
|
||||
|
||||
### Create a client secret
|
||||
|
||||
On the app's page, generate a client secret. You do not need a private key: DocsGPT only uses user tokens.
|
||||
|
||||
### Set the server settings
|
||||
|
||||
```env
|
||||
GITHUB_CLIENT_ID=Iv23li... # the app's Client ID
|
||||
GITHUB_CLIENT_SECRET=your-client-secret
|
||||
GITHUB_APP_SLUG=docsgpt-acme # from the app's public link, github.com/apps/<slug>
|
||||
```
|
||||
|
||||
Restart the API and the worker. **Sign in with GitHub** appears next to the token option.
|
||||
|
||||
</Steps>
|
||||
|
||||
#### Letting agents make changes
|
||||
|
||||
Under **Let agents use GitHub**, **Also let agents make changes (issues, comments, pull requests)** is off by default. Turned on, the tool uses GitHub's full endpoint `https://api.githubcopilot.com/mcp/` instead, which adds write actions such as creating issues, commenting and opening pull requests. The same switch is on the connection's page, on its GitHub tool; switching re-reads the actions from the other endpoint and keeps the permissions and fixed values of actions both endpoints have. Either way the tool sends the token only to `api.githubcopilot.com`.
|
||||
|
||||
An action is a **Write** unless GitHub marks it read-only, and writes default to **Needs approval**, so each change asks first until you choose **Always allow** for it. As everywhere, an agent called with its API key takes a write action only if the owner allows it in **Access details**, and an admin can turn changes off for everyone (see [Admin setup](#admin-setup)).
|
||||
|
||||
GitHub still checks the token: changes need a token, or a GitHub App, with write access to what they touch:
|
||||
|
||||
| To let agents | Repository permission |
|
||||
| --- | --- |
|
||||
| Create and edit issues, comment on issues | **Issues**: Read and write |
|
||||
| Open, update and review pull requests, comment on them | **Pull requests**: Read and write |
|
||||
| Create or edit files, push commits, create branches | **Contents**: Read and write |
|
||||
|
||||
Leave **Contents** read-only unless agents should edit files. Without a permission, GitHub refuses that action and the agent reports the error.
|
||||
|
||||
#### Public repositories and GITHUB_ACCESS_TOKEN
|
||||
|
||||
The **GitHub** tile under **Upload & web** in **Add knowledge** still ingests a public repository from its URL, without an account. `GITHUB_ACCESS_TOKEN`, if set, is used only there and only for public repositories, to raise GitHub's rate limit: DocsGPT checks that the repository is public first, because that token belongs to the server and not to the member asking. For a private repository, the form links to the member's GitHub connection.
|
||||
|
||||
<Callout type="warning" emoji="⚠️">
|
||||
Earlier versions read any repository `GITHUB_ACCESS_TOKEN` could see, private ones included. Sources made that way from private repositories stop syncing; recreate them from a GitHub connection.
|
||||
</Callout>
|
||||
|
||||
### Amazon S3
|
||||
|
||||
S3 needs no server settings. Each member connects with an access key that can list and read the bucket (`s3:ListBucket`, `s3:GetObject`), then picks a bucket and optional path prefix. A custom endpoint URL connects S3-compatible storage such as MinIO or Cloudflare R2.
|
||||
|
||||
### MCP presets
|
||||
|
||||
Notion, Linear, Atlassian, Sentry, Asana and Stripe are remote MCP servers that support OAuth with dynamic client registration, so they need no server settings. Members sign in with their own account. The presets ship in `docsgpt/connectors/presets/mcp.yaml`.
|
||||
|
||||
### Linear
|
||||
|
||||
One Linear sign-in gives agents Linear's tools and syncs Linear into Knowledge. After signing in, members turn on **Sync into Knowledge** and choose what to sync, or leave it off to keep only the tools; a connection's page adds more with **Sync more content**.
|
||||
|
||||
- **Teams and projects.** Each picked team or project brings its issues. An issue in both is synced once.
|
||||
- **Include comments** (on by default) adds each issue's comments to it.
|
||||
- **Include project documents** syncs the Linear documents of the picked projects.
|
||||
|
||||
Each issue becomes one document: its identifier and title, state, assignee, priority, labels, description and comments. It's filed under its team's key (`ENG/ENG-123.md`) and answers cite it with its Linear link. Archived issues are left out.
|
||||
|
||||
Linear's MCP server issues its own sign-in tokens, which work only with that server, so DocsGPT reads Linear through the same MCP tools the agents use. Nobody registers a Linear OAuth app. Tokens are renewed automatically. If renewing fails, the connection is marked **Reconnect needed** and its sources pause until you sign in again.
|
||||
|
||||
Each sync reads the source again in full, up to 500 issues and 100 documents, most recently updated first. For a large workspace, pick teams or projects rather than everything, and sync daily or weekly.
|
||||
|
||||
<Callout type="info">
|
||||
Background sync runs on the Celery worker and beat. Keep both running, as in the bundled Compose and Kubernetes files.
|
||||
</Callout>
|
||||
@@ -10,6 +10,10 @@ import { Steps } from 'nextra/components'
|
||||
|
||||
Connect your Confluence Cloud workspace to upload and process pages directly as an external knowledge base. Supports page content and attachments (PDFs, Office files, text files, images, and more). Authentication is handled via Atlassian OAuth 2.0 with automatic token refresh.
|
||||
|
||||
<Callout type="info">
|
||||
Members connect this service from **Settings > Connectors**, and one sign-in serves every source that uses it. Admins can check what is still missing in **Admin > Connectors**. See [Connectors](/Guides/Connectors).
|
||||
</Callout>
|
||||
|
||||
## Setup
|
||||
|
||||
<Steps>
|
||||
@@ -18,8 +22,8 @@ Connect your Confluence Cloud workspace to upload and process pages directly as
|
||||
|
||||
1. Go to [developer.atlassian.com/console/myapps](https://developer.atlassian.com/console/myapps/) and click **Create** > **OAuth 2.0 integration**
|
||||
2. Under **Authorization**, add a callback URL:
|
||||
- Local: `http://localhost:7091/api/connectors/callback?provider=confluence`
|
||||
- Production: `https://yourdomain.com/api/connectors/callback?provider=confluence`
|
||||
- Local: `http://127.0.0.1:7091/api/connectors/callback` (the default `CONNECTOR_REDIRECT_BASE_URI`; Admin > Connectors shows the exact value to copy)
|
||||
- Production: `https://yourdomain.com/api/connectors/callback` (the value of `CONNECTOR_REDIRECT_BASE_URI`, registered as-is)
|
||||
|
||||
### Step 2: Configure Permissions
|
||||
|
||||
@@ -56,7 +60,7 @@ VITE_CONFLUENCE_CLIENT_ID=your-atlassian-client-id
|
||||
|
||||
### Step 5: Restart and Use
|
||||
|
||||
Restart your application, then go to the upload section in DocsGPT and select **Confluence** as the source. You'll be redirected to Atlassian to sign in, then can browse spaces and select pages to process.
|
||||
Restart your application, then go to **Settings > Connectors** and pick **Confluence**. You'll be redirected to Atlassian to sign in, then can browse spaces and select pages to process.
|
||||
|
||||
</Steps>
|
||||
|
||||
@@ -64,6 +68,6 @@ Restart your application, then go to the upload section in DocsGPT and select **
|
||||
|
||||
- **Option not appearing** — Verify `VITE_CONFLUENCE_CLIENT_ID` is set in the frontend `.env`, then restart.
|
||||
- **Sign-in popup closes but the account never connects** — The frontend origin is not allowed to receive the result. Add it to `CONNECTOR_ALLOWED_ORIGINS` in the backend `.env`.
|
||||
- **Authentication failed** — Check that the callback URL matches exactly, including `?provider=confluence`.
|
||||
- **Authentication failed** — Check that the callback URL matches exactly and equals `CONNECTOR_REDIRECT_BASE_URI`, with no query parameters.
|
||||
- **No accessible sites** — Ensure the authenticating user has access to at least one Confluence Cloud site.
|
||||
- **Permission denied** — Verify that the Confluence API scopes are enabled in your Atlassian app settings.
|
||||
@@ -10,6 +10,10 @@ import { Steps } from 'nextra/components'
|
||||
|
||||
Connect your Google Drive account to upload and process files directly as an external knowledge base. Supports Google Workspace files (Docs, Sheets, Slides), Office files, PDFs, text files, CSVs, images, and more. Authentication is handled via Google OAuth 2.0 with automatic token refresh.
|
||||
|
||||
<Callout type="info">
|
||||
Members connect this service from **Settings > Connectors**, and one sign-in serves every source that uses it. Admins can check what is still missing in **Admin > Connectors**. See [Connectors](/Guides/Connectors).
|
||||
</Callout>
|
||||
|
||||
## Setup
|
||||
|
||||
<Steps>
|
||||
@@ -26,8 +30,8 @@ Connect your Google Drive account to upload and process files directly as an ext
|
||||
3. Select **Web application** as the application type
|
||||
4. Add your DocsGPT URL to **Authorized JavaScript origins** (e.g. `http://localhost:3000`)
|
||||
5. Add your callback URL to **Authorized redirect URIs**:
|
||||
- Local: `http://localhost:7091/api/connectors/callback?provider=google_drive`
|
||||
- Production: `https://yourdomain.com/api/connectors/callback?provider=google_drive`
|
||||
- Local: `http://127.0.0.1:7091/api/connectors/callback` (the default `CONNECTOR_REDIRECT_BASE_URI`; Admin > Connectors shows the exact value to copy)
|
||||
- Production: `https://yourdomain.com/api/connectors/callback` (the value of `CONNECTOR_REDIRECT_BASE_URI`, registered as-is)
|
||||
6. Click **Create** and copy the **Client ID** and **Client Secret**
|
||||
|
||||
### Step 3: Configure Environment Variables
|
||||
@@ -39,7 +43,7 @@ GOOGLE_CLIENT_ID=your-google-client-id
|
||||
GOOGLE_CLIENT_SECRET=your-google-client-secret
|
||||
```
|
||||
|
||||
Add to your frontend `.env` file:
|
||||
Optionally, to use Google's own file picker instead of DocsGPT's, add to your frontend `.env` file:
|
||||
|
||||
```env
|
||||
VITE_GOOGLE_CLIENT_ID=your-google-client-id
|
||||
@@ -49,23 +53,23 @@ VITE_GOOGLE_CLIENT_ID=your-google-client-id
|
||||
|----------|-------------|----------|
|
||||
| `GOOGLE_CLIENT_ID` | OAuth Client ID from GCP Credentials | Yes |
|
||||
| `GOOGLE_CLIENT_SECRET` | OAuth Client Secret from GCP Credentials | Yes |
|
||||
| `VITE_GOOGLE_CLIENT_ID` | Same Client ID, used by the frontend to show the Google Drive option | Yes |
|
||||
| `VITE_GOOGLE_CLIENT_ID` | Same Client ID, used by the frontend for Google's file picker | No |
|
||||
| `CONNECTOR_ALLOWED_ORIGINS` | Comma-separated frontend origins allowed to receive the sign-in result, e.g. `https://docsgpt.example.com`. Not needed when the frontend shares the API origin, or in local dev when the callback is on `localhost`/`127.0.0.1` and the frontend runs on port 5173 | When the frontend is on its own origin |
|
||||
|
||||
<Callout type="warning" emoji="⚠️">
|
||||
Make sure to use the same Google Client ID in both backend and frontend configurations.
|
||||
If you set `VITE_GOOGLE_CLIENT_ID`, use the same Client ID as the backend. Publish the OAuth consent screen (or use an internal Workspace app): apps left in Testing get refresh tokens that expire after seven days, which stops background sync.
|
||||
</Callout>
|
||||
|
||||
### Step 4: Restart and Use
|
||||
|
||||
Restart your application, then go to the upload section in DocsGPT and select **Google Drive** as the source. You'll be redirected to Google to sign in, then can browse and select files to process.
|
||||
Restart your application, then go to **Settings > Connectors** and pick **Google Drive**. You'll be redirected to Google to sign in, then can browse and select files to process.
|
||||
|
||||
</Steps>
|
||||
|
||||
## Troubleshooting
|
||||
|
||||
- **Option not appearing** — Verify `VITE_GOOGLE_CLIENT_ID` is set in the frontend `.env`, then restart.
|
||||
- **Authentication failed** — Check that the redirect URI matches exactly, including `?provider=google_drive`. Ensure the Google Drive API is enabled.
|
||||
- **Google Drive is not offered** — `GOOGLE_CLIENT_ID` or `GOOGLE_CLIENT_SECRET` is missing from the backend `.env`, so the connector stays off. Admin > Connectors lists which one.
|
||||
- **Authentication failed** — Check that the redirect URI matches exactly and equals `CONNECTOR_REDIRECT_BASE_URI`, with no query parameters. Ensure the Google Drive API is enabled.
|
||||
- **Sign-in popup closes but the account never connects** — The frontend origin is not allowed to receive the result. Add it to `CONNECTOR_ALLOWED_ORIGINS` in the backend `.env`.
|
||||
- **Permission denied** — Verify the OAuth consent screen is configured and the user has access to the target files.
|
||||
- **Files not processing** — Check backend logs and verify that backend environment variables are correctly set.
|
||||
|
||||
@@ -22,11 +22,11 @@ Only needed if your MCP servers use OAuth authentication:
|
||||
MCP_OAUTH_REDIRECT_URI=https://yourdomain.com/api/mcp_server/callback
|
||||
```
|
||||
|
||||
If not set, falls back to `API_URL/api/mcp_server/callback`.
|
||||
If not set, it is derived from the host of `CONNECTOR_REDIRECT_BASE_URI`, then from `API_URL`.
|
||||
|
||||
### Step 2: Add an MCP Server
|
||||
|
||||
Go to **Settings** > **Tools** > **Add Tool** > **MCP Server**. Enter the server URL, select an auth type, and click **Test Connection** to verify, then **Save**.
|
||||
Go to **Settings** > **Connectors** and pick a preset (Notion, Linear, Atlassian, Sentry, Asana, Stripe): **Sign in to Notion** opens the service's sign-in, and its tools are ready when you come back. For any other server, choose **Add custom connector** > **MCP server** and enter its URL and authentication; scopes and the timeout are under **Show advanced**. Enter the server URL, select an auth type, and click **Test Connection** to verify, then **Save**.
|
||||
|
||||
### Step 3: Enable for Your Agent
|
||||
|
||||
@@ -34,6 +34,8 @@ In your agent configuration, enable the MCP tools you want the agent to use.
|
||||
|
||||
</Steps>
|
||||
|
||||
Presets need no URL or form: they sign in with OAuth in one step. Admins can turn presets off one by one, and turn off custom MCP servers, in **Admin > Connectors**. See [Connectors](/Guides/Connectors).
|
||||
|
||||
## Authentication Types
|
||||
|
||||
| Auth Type | Config Fields |
|
||||
|
||||
@@ -10,6 +10,10 @@ import { Steps } from 'nextra/components'
|
||||
|
||||
Connect your SharePoint or OneDrive account to upload and process files directly as an external knowledge base. Supports Office files, PDFs, text files, CSVs, images, and more. Authentication is handled via Microsoft Entra ID (Azure AD) with automatic token refresh.
|
||||
|
||||
<Callout type="info">
|
||||
Members connect this service from **Settings > Connectors**, and one sign-in serves every source that uses it. Admins can check what is still missing in **Admin > Connectors**. See [Connectors](/Guides/Connectors).
|
||||
</Callout>
|
||||
|
||||
## Setup
|
||||
|
||||
<Steps>
|
||||
@@ -18,8 +22,8 @@ Connect your SharePoint or OneDrive account to upload and process files directly
|
||||
|
||||
1. Go to the [Azure Portal](https://portal.azure.com/) > **Microsoft Entra ID** > **App registrations** > **New registration**
|
||||
2. Set **Redirect URI** (Web) to:
|
||||
- Local: `http://localhost:7091/api/connectors/callback?provider=share_point`
|
||||
- Production: `https://yourdomain.com/api/connectors/callback?provider=share_point`
|
||||
- Local: `http://127.0.0.1:7091/api/connectors/callback` (the default `CONNECTOR_REDIRECT_BASE_URI`; Admin > Connectors shows the exact value to copy)
|
||||
- Production: `https://yourdomain.com/api/connectors/callback` (the value of `CONNECTOR_REDIRECT_BASE_URI`, registered as-is)
|
||||
|
||||
### Step 2: Configure API Permissions
|
||||
|
||||
@@ -53,13 +57,13 @@ MICROSOFT_TENANT_ID=your-azure-ad-tenant-id
|
||||
|
||||
### Step 5: Restart and Use
|
||||
|
||||
Restart your application, then go to the upload section in DocsGPT and select **SharePoint / OneDrive** as the source. You'll be redirected to Microsoft to sign in, then can browse and select files to process.
|
||||
Restart your application, then go to **Settings > Connectors** and pick **SharePoint**. You'll be redirected to Microsoft to sign in, then can browse and select files to process.
|
||||
|
||||
</Steps>
|
||||
|
||||
## Troubleshooting
|
||||
|
||||
- **Option not appearing** — Verify `MICROSOFT_CLIENT_ID` and `MICROSOFT_CLIENT_SECRET` are set, then restart.
|
||||
- **Authentication failed** — Check that the redirect URI matches exactly, including `?provider=share_point`.
|
||||
- **Authentication failed** — Check that the redirect URI matches exactly and equals `CONNECTOR_REDIRECT_BASE_URI`, with no query parameters.
|
||||
- **Sign-in popup closes but the account never connects** — The frontend origin is not allowed to receive the result. Add it to `CONNECTOR_ALLOWED_ORIGINS` in the backend `.env`.
|
||||
- **Permission denied** — Ensure admin consent is granted and the user has access to the target files.
|
||||
@@ -1,4 +1,8 @@
|
||||
export default {
|
||||
"Connectors": {
|
||||
"title": "🔌 Connectors",
|
||||
"href": "/Guides/Connectors"
|
||||
},
|
||||
"Customising-prompts": {
|
||||
"title": "️💻 Customising Prompts",
|
||||
"href": "/Guides/Customising-prompts"
|
||||
|
||||
@@ -90,7 +90,29 @@ PUT /api/sources/<source_id>/wiki/page # create or overwrite a page (
|
||||
|
||||
Human edits are stamped with `human` provenance and trigger the same re-embed as agent edits. Read access follows source sharing (owner or anyone the source is shared with); writing requires owner or team `editor` access.
|
||||
|
||||
## Edits from the API, widget and public links
|
||||
|
||||
An agent called with its API key (the website widget, the API) runs as its owner, so it could change any wiki the owner can edit. By default it can't: API and widget users can still read the wiki through the agent, but the agent isn't offered the create, edit, delete or rename actions in those chats, and the Wiki tool refuses them if it's asked anyway.
|
||||
|
||||
The wiki's owner can allow such edits. On the **Sources** page, open the wiki's menu, choose **Wiki settings**, and turn on **Let API and widget users edit this wiki**. The change applies from the next message, including in chats that are already open. Only the owner sees it; team editors can't change it. You and the wiki's editors can always edit it in DocsGPT, and so can agents you or they use there.
|
||||
|
||||
The switch doesn't cover public links. People who open an agent from its public link run it as themselves, so they can only edit wikis they could edit anyway (their own, or ones shared with them as an editor). Because the agent's prompt and sources belong to someone else, each of their wiki edits waits for their approval in the chat before it runs. A research agent can't stop to ask, so in a public-link chat it reads the wiki but doesn't edit it.
|
||||
|
||||
Scheduled and webhook runs don't get the Wiki tool, so they read a wiki through search but never edit it.
|
||||
|
||||
<Callout type="warning" emoji="⚠️">
|
||||
With authentication off (`AUTH_TYPE=None`, the default), every request, including one from the widget or the API, runs as the same `local` user who owns the agents. The switch then can't tell outside callers from you and has no effect, so turn authentication on if others can reach the widget or the API.
|
||||
</Callout>
|
||||
|
||||
From the API, only a signed-in session can change the setting; a personal access token can read it but not change it:
|
||||
|
||||
```text
|
||||
GET /api/sources/<source_id>/wiki/settings # {"allow_outside_edits": false, ...}
|
||||
PUT /api/sources/<source_id>/wiki/settings # owner only; body {"allow_outside_edits": true}
|
||||
```
|
||||
|
||||
## Related
|
||||
|
||||
- [Per-Source Configuration](/Sources/Per-source-configuration) — exposure and retrieval settings a wiki uses.
|
||||
- [Access Control & Teams](/Deploying/Access-Control) — sharing a wiki with a team.
|
||||
- [Access Control & Teams](/Deploying/Access-Control) — sharing a wiki with a team, and [what an agent's users reach](/Deploying/Access-Control#sharing-agents-and-what-they-use) when you share the agent.
|
||||
- [Connectors](/Guides/Connectors#agents-used-through-an-api-key) — what else an agent can and can't do for API, widget and public-link users.
|
||||
@@ -1,4 +1,8 @@
|
||||
export default {
|
||||
"Connectors": {
|
||||
"title": "🔌 Synced Sources (Connectors)",
|
||||
"href": "/Guides/Connectors"
|
||||
},
|
||||
"Per-source-configuration": {
|
||||
"title": "🎛️ Per-Source Configuration",
|
||||
"href": "/Sources/Per-source-configuration"
|
||||
|
||||
@@ -59,7 +59,7 @@ DocsGPT includes a suite of pre-built tools designed to expand its capabilities
|
||||
{
|
||||
title: 'Telegram Bot',
|
||||
link: 'https://github.com/arc53/DocsGPT/blob/main/docsgpt/agents/tools/telegram.py',
|
||||
description: 'Allows DocsGPT to send messages or images to Telegram chats via a Telegram Bot. Requires a bot token and chat ID.'
|
||||
description: 'Allows DocsGPT to send messages or images to Telegram chats via a Telegram Bot. Requires a bot token; an optional default chat ID sends every message to one chat.'
|
||||
},
|
||||
{
|
||||
title: 'PostgreSQL Database',
|
||||
@@ -121,7 +121,8 @@ Interacting with tools in DocsGPT is designed to be intuitive:
|
||||
|
||||
2. **Configuration in UI:**
|
||||
* Tools are generally managed and configured within the DocsGPT application's settings, found under a "Tools" section in the GUI.
|
||||
* For tools that interact with external services (like Brave Search, Telegram, or any service via the API Tool), you might need to provide authentication credentials (e.g., API keys, tokens) or specific endpoint information during the tool's setup in the UI.
|
||||
* Tools for external services (Brave Search, Telegram, ntfy, PostgreSQL, MCP servers) are set up from **Settings > Connectors**. You connect the service once; its credentials stay encrypted on the server and never reach the browser, and the tools it provides are grouped under that connection. Each action is marked **Read** or **Write** and can be set to **Always allow**, **Needs approval** or **Off**. See [Connectors](/Guides/Connectors).
|
||||
* When a tool's connection is missing or expired, the chat shows a **Connect** prompt and continues once you connect.
|
||||
|
||||
3. **Prompt Engineering for Tools:** While the LLM aims to intelligently use tools, for more complex or reliable agent-like behaviors, you might need to customize the system prompts. Modifying the prompt can guide the LLM on when and how to prioritize or chain tools to achieve specific outcomes, especially if you're building an agent designed to perform a certain sequence of actions every time. For more on this, see [Customising Prompts](/Guides/Customising-prompts).
|
||||
|
||||
|
||||
@@ -11,6 +11,23 @@ import { Callout } from 'nextra/components'
|
||||
**Upgrading from 0.16.x?** User data moved from MongoDB to Postgres in 0.17.0. Follow the [Postgres Migration guide](/Deploying/Postgres-Migration) before running `docker compose pull` or `git pull` — existing deployments will not start cleanly without it.
|
||||
</Callout>
|
||||
|
||||
## Connectors: set ENCRYPTION_SECRET_KEY first
|
||||
|
||||
Service credentials (OAuth tokens for Google Drive, SharePoint, Confluence and MCP servers, and API keys for tools) now live on **connections** and are encrypted with a key derived from `ENCRYPTION_SECRET_KEY`. Migration `0040_connections` runs on startup, encrypts the stored tokens and removes their plaintext copies.
|
||||
|
||||
<Callout type="warning">
|
||||
**Multi-user installs** (any `AUTH_TYPE`): set `ENCRYPTION_SECRET_KEY` to your own value **before** you upgrade, in the environment of the API, the worker and anything that runs migrations. The migration encrypts with the key it sees; changing it afterwards makes every connection ask its owner to reconnect. While the key is still the public default, DocsGPT refuses to store new credentials and Admin > Connectors shows a warning.
|
||||
</Callout>
|
||||
|
||||
If you already used a key and want to change it, see [rotating the key](/Guides/Connectors#encryption-key). Single-user local installs keep working with the default and log a warning at startup.
|
||||
|
||||
Other changes:
|
||||
|
||||
- Register `CONNECTOR_REDIRECT_BASE_URI` exactly as set (no `?provider=` query) as the redirect URI of each OAuth app. Admin > Connectors shows it.
|
||||
- `VITE_GOOGLE_CLIENT_ID` is optional now; without it, members pick Drive files in DocsGPT's own picker.
|
||||
- Synced sources run as their connection, without a browser session. Keep Celery beat running for scheduled syncs.
|
||||
- Downgrading to `0037_request_traces` decrypts the tokens back and gives each tool its key again.
|
||||
|
||||
## pip installs: data home moved
|
||||
|
||||
An installed package (`pip install docsgpt`, pipx, `uv tool`) used to keep its data home, meaning `.env`, `inputs/`, `indexes/` and `models/`, in the directory you ran `docsgpt api` and `docsgpt worker` from. It is now `~/.docsgpt/server` (`/opt/docsgpt` for root on Linux). Either move those files there, or set `DOCSGPT_HOME` to the old directory in the environment of both commands. Both commands print the data home they use, and point out a `.env` in the working directory that they no longer read. Source checkouts and the Docker images are not affected.
|
||||
|
||||
@@ -15,15 +15,17 @@ from docsgpt.api.answer.services.prompt_renderer import (
|
||||
resolve_prompt_skeleton,
|
||||
)
|
||||
from docsgpt.api.answer.services.stream_processor import (
|
||||
authorized_agent_sources,
|
||||
authorized_prompt_id,
|
||||
get_prompt,
|
||||
)
|
||||
from docsgpt.api.user.resource_access import ref_principal
|
||||
from docsgpt.core.settings import settings
|
||||
from docsgpt.guardrails.config import AgentConfig
|
||||
from docsgpt.quotas.service import QuotaExceededError, QuotaService
|
||||
from docsgpt.retriever.dispatcher import build_dispatcher
|
||||
from docsgpt.retriever.retriever_creator import RetrieverCreator
|
||||
from docsgpt.storage.db.repositories.sources import SourcesRepository
|
||||
from docsgpt.storage.db.session import db_readonly
|
||||
from docsgpt.storage.db.source_config import SourceConfig
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -74,6 +76,8 @@ def run_agent_headless(
|
||||
endpoint: str = "headless",
|
||||
chat_history: Optional[List[Dict[str, Any]]] = None,
|
||||
conversation_id: Optional[str] = None,
|
||||
external_caller: bool = False,
|
||||
public_link_caller: bool = False,
|
||||
request_id: Optional[str] = None,
|
||||
trace_user_id: Optional[str] = None,
|
||||
) -> Dict[str, Any]:
|
||||
@@ -86,6 +90,11 @@ def run_agent_headless(
|
||||
``trace_user_id`` owns the trace when the run belongs to someone other
|
||||
than the agent's owner (a schedule a user set on a shared agent), so the
|
||||
trace is visible wherever that user sees the run; it defaults to the owner.
|
||||
``external_caller`` (a schedule set through the API) and
|
||||
``public_link_caller`` (a schedule a public-link user set) mark a run for
|
||||
someone who can't approve for the owner: writes on the owner's accounts
|
||||
and credentials then run only when the agent's API write allowlist has
|
||||
them.
|
||||
|
||||
Raises:
|
||||
QuotaExceededError: If the agent owner's usage quota is exhausted.
|
||||
@@ -108,6 +117,8 @@ def run_agent_headless(
|
||||
endpoint=endpoint,
|
||||
chat_history=chat_history,
|
||||
conversation_id=conversation_id,
|
||||
external_caller=external_caller,
|
||||
public_link_caller=public_link_caller,
|
||||
)
|
||||
if outcome.get("error"):
|
||||
status = tracing.STATUS_ERROR
|
||||
@@ -128,6 +139,8 @@ def _run_agent_headless(
|
||||
endpoint: str = "headless",
|
||||
chat_history: Optional[List[Dict[str, Any]]] = None,
|
||||
conversation_id: Optional[str] = None,
|
||||
external_caller: bool = False,
|
||||
public_link_caller: bool = False,
|
||||
) -> Dict[str, Any]:
|
||||
from docsgpt.core.model_utils import (
|
||||
get_api_key_for_provider,
|
||||
@@ -148,20 +161,22 @@ def _run_agent_headless(
|
||||
raise QuotaExceededError(exceeded)
|
||||
|
||||
retriever_kind = agent_config.get("retriever", "classic")
|
||||
source_id = agent_config.get("source_id") or agent_config.get("source")
|
||||
source_active: Any = {}
|
||||
if source_id:
|
||||
# Every source a chat with this agent searches: the primary and the
|
||||
# extras, each owned or team-shared to the owner, else attached by an
|
||||
# editor who still qualifies.
|
||||
sources_row = dict(agent_config)
|
||||
sources_row["source_id"] = agent_config.get("source_id") or agent_config.get("source")
|
||||
primary, source_docs = None, []
|
||||
if sources_row["source_id"] or sources_row.get("extra_source_ids"):
|
||||
with db_readonly() as conn:
|
||||
# Owned or team-shared to the owner, else attached by an editor
|
||||
# who still qualifies; read unscoped once authorized.
|
||||
src_row = (
|
||||
SourcesRepository(conn).get_by_id(str(source_id))
|
||||
if ref_principal(conn, "agent", agent_config, "source", str(source_id))
|
||||
else None
|
||||
)
|
||||
if src_row:
|
||||
source_active = str(src_row["id"])
|
||||
retriever_kind = src_row.get("retriever", retriever_kind)
|
||||
primary, source_docs = authorized_agent_sources(conn, sources_row)
|
||||
if primary:
|
||||
retriever_kind = primary.get("retriever") or retriever_kind
|
||||
per_source = [
|
||||
{"id": str(doc["id"]), "retrieval": SourceConfig.parse(doc.get("config")).retrieval}
|
||||
for doc in source_docs
|
||||
]
|
||||
source_active: Any = [entry["id"] for entry in per_source] or {}
|
||||
source = {"active_docs": source_active}
|
||||
# ``chunks=0`` switches retrieval off; only a missing value takes the default.
|
||||
raw_chunks = agent_config.get("chunks")
|
||||
@@ -196,8 +211,7 @@ def _run_agent_headless(
|
||||
system_api_key = get_api_key_for_provider(provider or settings.LLM_PROVIDER)
|
||||
doc_token_limit = calculate_doc_token_budget(model_id=model_id, user_id=owner)
|
||||
|
||||
retriever = RetrieverCreator.create_retriever(
|
||||
retriever_kind,
|
||||
retriever_kwargs: Dict[str, Any] = dict(
|
||||
source=source,
|
||||
chat_history=chat_history or [],
|
||||
prompt=prompt,
|
||||
@@ -208,6 +222,13 @@ def _run_agent_headless(
|
||||
agent_id=agent_id,
|
||||
decoded_token=decoded_token,
|
||||
)
|
||||
# Routed per source like a chat's pre-fetch, so each source keeps its own
|
||||
# retriever and retrieval settings.
|
||||
retriever = build_dispatcher(
|
||||
lambda: RetrieverCreator.create_retriever(retriever_kind, **retriever_kwargs),
|
||||
sources=per_source,
|
||||
**retriever_kwargs,
|
||||
)
|
||||
retrieved_docs: List[Dict[str, Any]] = []
|
||||
try:
|
||||
docs = retriever.search(query)
|
||||
@@ -223,6 +244,9 @@ def _run_agent_headless(
|
||||
agent_id=agent_id,
|
||||
headless=True,
|
||||
tool_allowlist=list(tool_allowlist or []),
|
||||
external_caller=external_caller,
|
||||
public_link_caller=public_link_caller,
|
||||
api_write_allowlist=AgentConfig.parse(agent_config.get("config")).api_write_allowlist,
|
||||
)
|
||||
if conversation_id:
|
||||
tool_executor.conversation_id = str(conversation_id)
|
||||
|
||||
@@ -6,7 +6,7 @@ from typing import Any, Dict, Generator, List, Optional
|
||||
|
||||
from docsgpt import tracing
|
||||
from docsgpt.agents.base import BaseAgent
|
||||
from docsgpt.agents.tool_executor import ToolExecutor
|
||||
from docsgpt.agents.tool_executor import ToolExecutor, journal_refused_call
|
||||
from docsgpt.agents.tools.graph_search import add_graph_search_tool
|
||||
from docsgpt.agents.tools.internal_search import add_internal_search_tool
|
||||
from docsgpt.agents.tools.wiki import add_wiki_tool
|
||||
@@ -582,24 +582,31 @@ class ResearchAgent(BaseAgent):
|
||||
search_returned_empty = False
|
||||
|
||||
for call in tool_calls:
|
||||
gen = executor.execute(
|
||||
tools_dict, call, self.llm.__class__.__name__
|
||||
)
|
||||
result = None
|
||||
call_id = None
|
||||
while True:
|
||||
try:
|
||||
event = next(gen)
|
||||
# Log tool_call status events instead of discarding them
|
||||
if isinstance(event, dict) and event.get("type") == "tool_call":
|
||||
logger.debug(
|
||||
"Tool %s status: %s",
|
||||
event.get("data", {}).get("action_name", ""),
|
||||
event.get("data", {}).get("status", ""),
|
||||
)
|
||||
except StopIteration as e:
|
||||
result, call_id = e.value
|
||||
break
|
||||
# A step runs inside one turn and nobody can answer a pause here,
|
||||
# so a call that would pause (approval, a connection, the client,
|
||||
# an outside caller's write on the owner's account) is refused.
|
||||
refusal = self._refuse_paused_call(tools_dict, call, executor)
|
||||
if refusal is not None:
|
||||
result, call_id = refusal
|
||||
else:
|
||||
gen = executor.execute(
|
||||
tools_dict, call, self.llm.__class__.__name__
|
||||
)
|
||||
result = None
|
||||
call_id = None
|
||||
while True:
|
||||
try:
|
||||
event = next(gen)
|
||||
# Log tool_call status events instead of discarding them
|
||||
if isinstance(event, dict) and event.get("type") == "tool_call":
|
||||
logger.debug(
|
||||
"Tool %s status: %s",
|
||||
event.get("data", {}).get("action_name", ""),
|
||||
event.get("data", {}).get("status", ""),
|
||||
)
|
||||
except StopIteration as e:
|
||||
result, call_id = e.value
|
||||
break
|
||||
|
||||
# Detect empty search results for refinement
|
||||
is_search = "search" in (call.name or "").lower()
|
||||
@@ -636,6 +643,60 @@ class ResearchAgent(BaseAgent):
|
||||
|
||||
return messages, search_returned_empty
|
||||
|
||||
def _refuse_paused_call(
|
||||
self, tools_dict: Dict, call, executor: ToolExecutor
|
||||
) -> Optional[tuple[str, str]]:
|
||||
"""Refuse a call ``check_pause`` would pause on, as a headless run does.
|
||||
|
||||
Args:
|
||||
tools_dict: The step's tools.
|
||||
call: The model's tool call.
|
||||
executor: The run's executor.
|
||||
|
||||
Returns:
|
||||
``(tool result, call id)`` when the call is refused, else None.
|
||||
"""
|
||||
pause_info = executor.check_pause(
|
||||
tools_dict, call, self.llm.__class__.__name__
|
||||
)
|
||||
if not pause_info:
|
||||
return None
|
||||
pause_type = pause_info.get("pause_type")
|
||||
if pause_type == "headless_denied":
|
||||
reason = pause_info.get("deny_reason") or "This tool can't run here."
|
||||
result = f"Tool denied: {reason}"
|
||||
journal_error = f"headless: {reason}" if executor.headless else f"denied: {reason}"
|
||||
if executor.headless:
|
||||
executor.headless_denials.append(pause_info)
|
||||
elif pause_info.get("connection_required"):
|
||||
result = (
|
||||
"Tool not run: its service needs to be connected first, and a "
|
||||
"research step can't wait for that."
|
||||
)
|
||||
journal_error = "research: connection required"
|
||||
elif pause_type == "requires_client_execution":
|
||||
result = (
|
||||
"Tool not run: it runs in the user's app, which a research step "
|
||||
"can't reach."
|
||||
)
|
||||
journal_error = "research: client-side tool"
|
||||
else:
|
||||
result = (
|
||||
"Tool not run: this action needs the user's approval, which a "
|
||||
"research step can't ask for. Tell the user it needs their "
|
||||
"approval in a regular chat."
|
||||
)
|
||||
journal_error = "research: approval required"
|
||||
logger.info(
|
||||
"research_step_tool_refused",
|
||||
extra={
|
||||
"action_name": pause_info.get("action_name"),
|
||||
"pause_type": pause_type,
|
||||
},
|
||||
)
|
||||
journal_refused_call(executor, pause_info, journal_error)
|
||||
return result, pause_info["call_id"]
|
||||
|
||||
def _collect_step_sources(self):
|
||||
"""Register the search tools' docs (internal search and graph pages) with CitationManager."""
|
||||
for doc in self._search_tool_docs():
|
||||
|
||||
+480
-53
@@ -15,6 +15,7 @@ from docsgpt.agents.default_tools import (
|
||||
synthesized_default_tools,
|
||||
)
|
||||
from docsgpt import tracing
|
||||
from docsgpt.agents.tool_pins import iter_parameters, llm_fills, resolve_arguments, sent_arguments
|
||||
from docsgpt.agents.tools.tool_action_parser import ToolActionParser
|
||||
from docsgpt.agents.tools.tool_manager import ToolManager
|
||||
from docsgpt.guardrails.types import Stage as GuardrailStage, resolve_tool_result
|
||||
@@ -188,6 +189,16 @@ def _requires_approval(tool: Dict, action: Dict) -> bool:
|
||||
return bool((tool.get("config") or {}).get("require_approval"))
|
||||
|
||||
|
||||
def _account_slug(account: Optional[str], limit: int = 24) -> str:
|
||||
"""An account name as a function-name suffix: ``Ops: on-call!`` → ``ops_on_call``.
|
||||
|
||||
Empty when nothing ASCII is left (a name in another script); the caller
|
||||
then numbers the duplicates instead.
|
||||
"""
|
||||
slug = re.sub(r"[^a-z0-9]+", "_", str(account or "").lower()).strip("_")
|
||||
return slug[:limit].rstrip("_")
|
||||
|
||||
|
||||
def _sanitize_tool_prefix(tool_name: Optional[str]) -> str:
|
||||
"""Reduce a tool name to characters allowed in function-call names."""
|
||||
return re.sub(r"[^a-zA-Z0-9_-]+", "_", str(tool_name or "")).strip("_")
|
||||
@@ -472,6 +483,36 @@ def _mark_failed(
|
||||
logger.exception("tool_call_attempts failed-write failed for %s", call_id)
|
||||
|
||||
|
||||
def journal_refused_call(executor: Any, pause_info: Dict, error: str) -> None:
|
||||
"""Journal a tool call that was refused instead of paused, as failed.
|
||||
|
||||
A headless run and a research step can't pause for anyone, so a call
|
||||
``check_pause`` would pause on is answered with a refusal. Journaling it
|
||||
keeps the refusal visible to the reconciler and tool analytics.
|
||||
|
||||
Args:
|
||||
executor: The run's ``ToolExecutor``.
|
||||
pause_info: What ``check_pause`` returned for the call.
|
||||
error: The failure recorded on the journal row.
|
||||
"""
|
||||
if _record_proposed(
|
||||
pause_info["call_id"],
|
||||
pause_info["tool_name"],
|
||||
pause_info["action_name"],
|
||||
pause_info.get("arguments") or {},
|
||||
tool_id=pause_info.get("tool_id"),
|
||||
message_id=getattr(executor, "message_id", None),
|
||||
user_id=getattr(executor, "user", None),
|
||||
agent_id=getattr(executor, "agent_id", None),
|
||||
):
|
||||
_mark_failed(
|
||||
pause_info["call_id"],
|
||||
error,
|
||||
message_id=getattr(executor, "message_id", None),
|
||||
user_id=getattr(executor, "user", None),
|
||||
)
|
||||
|
||||
|
||||
class ToolExecutor:
|
||||
"""Handles tool discovery, preparation, and execution.
|
||||
|
||||
@@ -487,6 +528,9 @@ class ToolExecutor:
|
||||
*,
|
||||
headless: bool = False,
|
||||
tool_allowlist: Optional[List[str]] = None,
|
||||
external_caller: bool = False,
|
||||
public_link_caller: bool = False,
|
||||
api_write_allowlist: Optional[List[str]] = None,
|
||||
):
|
||||
self.user_api_key = user_api_key
|
||||
self.user = user
|
||||
@@ -497,6 +541,15 @@ class ToolExecutor:
|
||||
self.headless = bool(headless)
|
||||
# Tool-instance ids pre-authorized for headless approval-gated execution.
|
||||
self.tool_allowlist: set = {str(x) for x in tool_allowlist} if tool_allowlist else set()
|
||||
# Someone calling the agent with its API key (widget, API): the run
|
||||
# uses the owner's accounts and nobody can approve, so writes on a
|
||||
# connected account run only when the owner allowlisted them.
|
||||
self.external_caller = bool(external_caller)
|
||||
# Someone who reaches the agent only through its public link: they
|
||||
# may not approve writes on the owner's account either, so those run
|
||||
# only when allowlisted. Their own account (member mode) is theirs.
|
||||
self.public_link_caller = bool(public_link_caller)
|
||||
self.api_write_allowlist: set = {str(x) for x in api_write_allowlist or []}
|
||||
# Set by BaseAgent._prepare_tools when the agent has tool-stage controls.
|
||||
self.guardrail_engine = None
|
||||
self.tool_calls: List[Dict] = []
|
||||
@@ -505,9 +558,15 @@ class ToolExecutor:
|
||||
# get_tools() resolves EXACTLY these ids — builtin synthetic ids and
|
||||
# user_tools rows alike — with no defaults mixed in. None = unscoped.
|
||||
self.allowed_tool_ids: Optional[List[str]] = None
|
||||
# Who an explicit tool-id scope resolves as: the workflow owner for a
|
||||
# node, so whoever runs it gets the owner's tools. None = ``user``.
|
||||
self.tool_owner: Optional[str] = None
|
||||
# Tool id -> the user to resolve it as, for a workflow node's tools
|
||||
# sponsored by an editor (see resource_access.active_sponsor).
|
||||
self.tool_principals: Dict[str, str] = {}
|
||||
# The workflow row a node's tools belong to, so a dropped one is
|
||||
# logged with why it doesn't run.
|
||||
self.tool_holder: Optional[Dict] = None
|
||||
self.conversation_id: Optional[str] = None
|
||||
# Set by the workflow engine for agent nodes so run-scoped tools
|
||||
# (artifact_generator / code_executor) address artifacts by the
|
||||
@@ -528,6 +587,11 @@ class ToolExecutor:
|
||||
self._tool_to_name: Dict[Tuple[str, str], str] = {}
|
||||
# Filled by the LLMHandler.handle_tool_calls headless loop.
|
||||
self.headless_denials: List[Dict] = []
|
||||
# Per-turn connection resolution for connection-backed tools, keyed
|
||||
# by tool row id, so check_pause and execute share one lookup.
|
||||
self._connections: Dict[str, Any] = {}
|
||||
# Tool parameters those connections set (Telegram's default chat).
|
||||
self._connection_params: Dict[str, Dict] = {}
|
||||
|
||||
def get_tools(self) -> Dict[str, Dict]:
|
||||
"""Load tool configs from DB based on user context.
|
||||
@@ -559,27 +623,45 @@ class ToolExecutor:
|
||||
"""Resolve an explicit tool-id scope — exactly these ids, no defaults.
|
||||
|
||||
Used by workflow agent nodes: the node's configured tools (builtin
|
||||
synthetic ids like Artifact/Code Executor/Read Document, or the user's
|
||||
``user_tools`` rows) are the node's WHOLE toolset. An unresolvable id
|
||||
is dropped with a warning rather than failing the node.
|
||||
synthetic ids like Artifact/Code Executor/Read Document, or the
|
||||
``user_tools`` rows of ``tool_owner``) are the node's WHOLE toolset.
|
||||
Rows resolve as the workflow owner, then as the editor who attached
|
||||
them, never as whoever runs the workflow — the same rule as an agent's
|
||||
own tools. An unresolvable id is dropped with a warning rather than
|
||||
failing the node.
|
||||
"""
|
||||
if not tool_ids:
|
||||
return {}
|
||||
principal = self.tool_owner or self.user
|
||||
with db_readonly() as conn:
|
||||
tools_repo = UserToolsRepository(conn)
|
||||
tools: List[Dict] = []
|
||||
for tid in tool_ids:
|
||||
row = resolve_tool_by_id(tid, self.user, user_tools_repo=tools_repo)
|
||||
row = resolve_tool_by_id(tid, principal, user_tools_repo=tools_repo)
|
||||
if row is None and str(tid) in self.tool_principals:
|
||||
row = resolve_tool_by_id(tid, self.tool_principals[str(tid)], user_tools_repo=tools_repo)
|
||||
if row is None:
|
||||
logger.warning("tool id %s did not resolve; dropped from scoped toolset", tid)
|
||||
self._log_dropped_scoped_tool(conn, tid, principal)
|
||||
continue
|
||||
if self.headless and is_headless_excluded_tool(row.get("name")):
|
||||
continue
|
||||
tools.append(row)
|
||||
return {str(tool["id"]): tool for tool in tools}
|
||||
|
||||
def _log_dropped_scoped_tool(self, conn, tool_id: str, principal: Optional[str]) -> None:
|
||||
"""Log a node tool the scoped toolset leaves out, with why when the workflow is known."""
|
||||
# Lazy: docsgpt.api's package import pulls in every route module.
|
||||
from docsgpt.api.user.resource_access import log_stopped, resolve_holder_tool
|
||||
|
||||
holder = {**(self.tool_holder or {}), "user_id": principal}
|
||||
reason = None
|
||||
try:
|
||||
_row, access = resolve_holder_tool(conn, "workflow", holder, tool_id)
|
||||
reason = access.reason
|
||||
except Exception:
|
||||
logger.exception("Could not tell why tool %s does not resolve", tool_id)
|
||||
log_stopped("workflow", holder, "tool", tool_id, reason)
|
||||
|
||||
def _get_tools_by_api_key(self, api_key: str) -> Dict[str, Dict]:
|
||||
"""Resolve an agent's toolset — exactly ``agents.tools``, no defaults."""
|
||||
# Per-operation session: the answer pipeline spans a long-lived
|
||||
@@ -595,13 +677,14 @@ class ToolExecutor:
|
||||
row = resolve_tool_by_id(tid, owner, user_tools_repo=tools_repo)
|
||||
if row is None:
|
||||
# A tool the owner can't use runs as the editor who
|
||||
# attached it, while they still qualify.
|
||||
# attached it, while they still qualify: the same check
|
||||
# the agent page's run state uses.
|
||||
# Lazy: docsgpt.api's package import pulls in every route module.
|
||||
from docsgpt.api.user.resource_access import active_sponsor
|
||||
from docsgpt.api.user.resource_access import log_stopped, resolve_holder_tool
|
||||
|
||||
sponsor = active_sponsor(conn, "agent", agent_data, "tool", str(tid))
|
||||
if sponsor:
|
||||
row = resolve_tool_by_id(tid, sponsor, user_tools_repo=tools_repo)
|
||||
row, access = resolve_holder_tool(conn, "agent", agent_data, tid, tools_repo=tools_repo)
|
||||
if row is None:
|
||||
log_stopped("agent", agent_data, "tool", tid, access.reason)
|
||||
if row is None:
|
||||
continue
|
||||
# Workflow-only builtins (read_document) never resolve for a
|
||||
@@ -791,15 +874,36 @@ class ToolExecutor:
|
||||
self._tool_to_name = {}
|
||||
all_llm_names: set = set()
|
||||
|
||||
# Connection tools that share an action name are told apart by what
|
||||
# they connect to: the service ("search" on Notion and on Linear), or
|
||||
# the account when one service is connected twice (two Telegram bots).
|
||||
connected: Dict[int, Tuple[str, str]] = {}
|
||||
for index, (tool_id, _tool_name, action_name, _action, is_client) in enumerate(entries):
|
||||
if name_counts[action_name] > 1 and not is_client:
|
||||
names = self._connection_names(tools_dict[tool_id])
|
||||
if names:
|
||||
connected[index] = names
|
||||
per_service = Counter((entries[i][2], service) for i, (service, _account) in connected.items())
|
||||
|
||||
result = []
|
||||
for tool_id, tool_name, action_name, action, is_client in entries:
|
||||
for index, (tool_id, tool_name, action_name, action, is_client) in enumerate(entries):
|
||||
service, account = connected.get(index, (None, None))
|
||||
# The account is named only where it is what tells tools apart.
|
||||
if service is not None and per_service[(action_name, service)] < 2:
|
||||
account = None
|
||||
slug = _account_slug(account or service)
|
||||
if name_counts[action_name] == 1 and len(action_name) <= _MAX_LLM_NAME_LEN:
|
||||
llm_name = action_name
|
||||
else:
|
||||
# An over-long unique name skips the prefix — it needs
|
||||
# truncation, not disambiguation.
|
||||
prefix = _sanitize_tool_prefix(tool_name) if name_counts[action_name] > 1 else ""
|
||||
base = f"{prefix}_{action_name}" if prefix and not action_name.startswith(f"{prefix}_") else action_name
|
||||
if slug:
|
||||
base = f"{action_name}_{slug}"
|
||||
elif prefix and not action_name.startswith(f"{prefix}_"):
|
||||
base = f"{prefix}_{action_name}"
|
||||
else:
|
||||
base = action_name
|
||||
base = base[:_MAX_LLM_NAME_LEN]
|
||||
# A duplicated bare name stays ambiguous, and a candidate
|
||||
# must not steal a unique action's name or one already taken.
|
||||
@@ -818,31 +922,54 @@ class ToolExecutor:
|
||||
if is_client:
|
||||
params = action.get("parameters", {})
|
||||
else:
|
||||
params = self._build_tool_parameters(action)
|
||||
params = self._build_tool_parameters(
|
||||
action, hidden=set(self._connection_parameters(tools_dict[tool_id])),
|
||||
)
|
||||
|
||||
description = action.get("description", "")
|
||||
if account:
|
||||
description = f"{description} ({service} account: {account})".strip()
|
||||
elif service:
|
||||
description = f"{description} ({service})".strip()
|
||||
result.append(
|
||||
{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": llm_name,
|
||||
"description": action.get("description", ""),
|
||||
"description": description,
|
||||
"parameters": params,
|
||||
},
|
||||
}
|
||||
)
|
||||
return result
|
||||
|
||||
def _build_tool_parameters(self, action: Dict) -> Dict:
|
||||
def _connection_names(self, tool_data: Dict) -> Optional[Tuple[str, str]]:
|
||||
"""``(service, account)`` a connection tool runs with, e.g. ``("Telegram", "Alerts bot")``."""
|
||||
if not tool_data.get("connection_id") or tool_data.get("client_side"):
|
||||
return None
|
||||
resolved = self._resolve_connection(tool_data)
|
||||
if resolved is None or resolved.row is None:
|
||||
return None
|
||||
from docsgpt.connectors.service import account_name
|
||||
|
||||
return resolved.connector_name or tool_data.get("name") or "", account_name(resolved.row)
|
||||
|
||||
def _build_tool_parameters(self, action: Dict, hidden: Optional[set] = None) -> Dict:
|
||||
"""The JSON schema the model sees for ``action``.
|
||||
|
||||
Parameters the model does not fill (fixed values) and ``hidden`` ones
|
||||
(values the connection fixes) are left out, so the model is never
|
||||
asked for them.
|
||||
"""
|
||||
params = {"type": "object", "properties": {}, "required": []}
|
||||
for param_type in ["query_params", "headers", "body", "parameters"]:
|
||||
if param_type in action and action[param_type].get("properties"):
|
||||
for k, v in action[param_type]["properties"].items():
|
||||
if v.get("filled_by_llm", True):
|
||||
params["properties"][k] = {
|
||||
key: value for key, value in v.items() if key not in ("filled_by_llm", "value", "required")
|
||||
}
|
||||
if v.get("required", False):
|
||||
params["required"].append(k)
|
||||
for _section, k, v in iter_parameters(action):
|
||||
if not llm_fills(v) or (hidden and k in hidden):
|
||||
continue
|
||||
params["properties"][k] = {
|
||||
key: value for key, value in v.items() if key not in ("filled_by_llm", "value", "required")
|
||||
}
|
||||
if v.get("required", False):
|
||||
params["required"].append(k)
|
||||
return params
|
||||
|
||||
def _guardrail_tool_result(self, result: Any, tool_name: str, action_name: str) -> Any:
|
||||
@@ -866,6 +993,70 @@ class ToolExecutor:
|
||||
return result
|
||||
return resolve_tool_result(result, decision)
|
||||
|
||||
def _resolve_connection(self, tool_data: Dict):
|
||||
"""The connection a connection-backed tool runs with this turn, or None."""
|
||||
if not tool_data.get("connection_id") or tool_data.get("client_side"):
|
||||
return None
|
||||
key = str(tool_data.get("id") or tool_data.get("connection_id"))
|
||||
if key not in self._connections:
|
||||
from docsgpt.connectors.resolve import resolve_connection
|
||||
|
||||
try:
|
||||
self._connections[key] = resolve_connection(tool_data, self.user)
|
||||
except Exception:
|
||||
logger.exception("connection resolution failed for tool %s", key)
|
||||
self._connections[key] = None
|
||||
return self._connections[key]
|
||||
|
||||
def _connection_parameters(self, tool_data: Dict) -> Dict:
|
||||
"""Tool parameters the connection a call runs with sets, e.g. Telegram's default chat.
|
||||
|
||||
Resolved like the credentials: in member mode each member's own
|
||||
connection (and so their own chat) applies. Only connectors with
|
||||
such fields are looked up.
|
||||
"""
|
||||
if not tool_data.get("connection_id") or tool_data.get("client_side"):
|
||||
return {}
|
||||
from docsgpt.connectors import catalog
|
||||
|
||||
definition = catalog.definition_for_tool(tool_data.get("name") or "")
|
||||
if definition is None or not catalog.parameter_fields(definition.key):
|
||||
return {}
|
||||
key = str(tool_data.get("id") or tool_data.get("connection_id"))
|
||||
if key not in self._connection_params:
|
||||
resolved = self._resolve_connection(tool_data)
|
||||
params: Dict = {}
|
||||
if resolved is not None and resolved.available:
|
||||
from docsgpt.connectors.service import connection_parameters
|
||||
|
||||
try:
|
||||
params = connection_parameters(resolved.row)
|
||||
except Exception:
|
||||
logger.exception("connection parameters failed for tool %s", key)
|
||||
self._connection_params[key] = params
|
||||
return self._connection_params[key]
|
||||
|
||||
@staticmethod
|
||||
def _connection_payload(resolved) -> Dict:
|
||||
"""What the chat's Connect card needs; never an account or a secret.
|
||||
|
||||
``connection_id`` is only the caller's own connection, which the card
|
||||
reconnects in place; an owner's account (``owner_account``) is not
|
||||
the caller's to reconnect.
|
||||
"""
|
||||
payload = {
|
||||
"connector_key": resolved.connector_key,
|
||||
"connector_name": resolved.connector_name,
|
||||
"status": (
|
||||
"missing" if resolved.row is None
|
||||
else (resolved.row.get("status") or "reconnect_needed")
|
||||
),
|
||||
"owner_account": bool(resolved.delegated),
|
||||
}
|
||||
if resolved.row is not None and not resolved.delegated:
|
||||
payload["connection_id"] = str(resolved.row["id"])
|
||||
return payload
|
||||
|
||||
def check_pause(self, tools_dict: Dict, call, llm_class_name: str) -> Optional[Dict]:
|
||||
"""Return a pending-action dict (approval / client / headless_denied) or None.
|
||||
|
||||
@@ -912,6 +1103,41 @@ class ToolExecutor:
|
||||
"thought_signature": getattr(call, "thought_signature", None),
|
||||
}
|
||||
|
||||
# A tool whose connection needs signing in pauses on a Connect card
|
||||
# (the approval card's connection variant) instead of failing; the
|
||||
# user connects, then continues, and the pending call resumes.
|
||||
resolved = self._resolve_connection(tool_data)
|
||||
if resolved is not None and not resolved.available:
|
||||
if self.headless or self.external_caller:
|
||||
return {
|
||||
"call_id": call_id,
|
||||
"name": llm_name,
|
||||
"tool_name": tool_data.get("name", "unknown"),
|
||||
"tool_id": tool_id,
|
||||
"action_name": action_name,
|
||||
"llm_name": llm_name,
|
||||
"arguments": arguments,
|
||||
"pause_type": "headless_denied",
|
||||
"deny_reason": (
|
||||
f"{resolved.connector_name or 'This service'} needs to be connected "
|
||||
"before this tool can run."
|
||||
),
|
||||
"error_type": "connection_required",
|
||||
"thought_signature": getattr(call, "thought_signature", None),
|
||||
}
|
||||
return {
|
||||
"call_id": call_id,
|
||||
"name": llm_name,
|
||||
"tool_name": tool_data.get("name", "unknown"),
|
||||
"tool_id": tool_id,
|
||||
"action_name": action_name,
|
||||
"llm_name": llm_name,
|
||||
"arguments": arguments,
|
||||
"pause_type": "awaiting_approval",
|
||||
"connection_required": self._connection_payload(resolved),
|
||||
"thought_signature": getattr(call, "thought_signature", None),
|
||||
}
|
||||
|
||||
# Approval required
|
||||
if tool_data["name"] == "api_tool":
|
||||
action_data = tool_data.get("config", {}).get("actions", {}).get(action_name, {})
|
||||
@@ -948,6 +1174,74 @@ class ToolExecutor:
|
||||
or require_approval
|
||||
)
|
||||
|
||||
# An admin forbade changes through this connector (GitHub): its tool
|
||||
# already calls the read-only endpoint, so say why instead of failing.
|
||||
if resolved is not None and not resolved.writes_allowed:
|
||||
from docsgpt.connectors.permissions import ACCESS_WRITE, action_access
|
||||
|
||||
if action_access(tool_data.get("name"), action_data) == ACCESS_WRITE:
|
||||
return {
|
||||
"call_id": call_id,
|
||||
"name": llm_name,
|
||||
"tool_name": tool_data.get("name", "unknown"),
|
||||
"tool_id": tool_id,
|
||||
"action_name": action_name,
|
||||
"llm_name": llm_name,
|
||||
"arguments": arguments,
|
||||
"pause_type": "headless_denied",
|
||||
"deny_reason": (
|
||||
f"An admin turned off changes through {resolved.connector_name or 'this service'}. "
|
||||
"It can only look things up."
|
||||
),
|
||||
"error_type": "tool_not_allowed",
|
||||
"thought_signature": getattr(call, "thought_signature", None),
|
||||
}
|
||||
|
||||
# A member running someone else's account (a shared tool in owner
|
||||
# mode) always confirms write actions, whatever the owner chose for
|
||||
# themselves.
|
||||
if not require_approval and resolved is not None and resolved.delegated:
|
||||
from docsgpt.connectors.permissions import ACCESS_WRITE, action_access
|
||||
|
||||
require_approval = action_access(tool_data.get("name"), action_data) == ACCESS_WRITE
|
||||
|
||||
# An API-key caller writes with the owner's credentials only with the
|
||||
# owner's say-so: nobody can approve in a widget, and "Always allow"
|
||||
# was the owner's choice for themselves, not for anyone with the key.
|
||||
# A public-link user is a stranger to the owner, so their approval
|
||||
# can't stand in for the owner's either.
|
||||
if (self.external_caller or self.public_link_caller) and self._on_owner_credentials(
|
||||
tool_data, resolved, action_name
|
||||
):
|
||||
from docsgpt.connectors.permissions import ACCESS_WRITE, action_access
|
||||
|
||||
if action_access(tool_data.get("name"), action_data) == ACCESS_WRITE:
|
||||
entry = f"{tool_data.get('id') or tool_id}:{action_name}"
|
||||
if entry in self.api_write_allowlist:
|
||||
return None
|
||||
route = "for API or widget callers" if self.external_caller else "from its public link"
|
||||
target = (
|
||||
f"the owner's {resolved.connector_name} account"
|
||||
if resolved is not None and resolved.connector_name
|
||||
else "the owner's credentials"
|
||||
)
|
||||
return {
|
||||
"call_id": call_id,
|
||||
"name": llm_name,
|
||||
"tool_name": tool_data.get("name", "unknown"),
|
||||
"tool_id": tool_id,
|
||||
"action_name": action_name,
|
||||
"llm_name": llm_name,
|
||||
"arguments": arguments,
|
||||
"pause_type": "headless_denied",
|
||||
"deny_reason": (
|
||||
f"This agent can't take this action with {target} {route}. "
|
||||
"The owner can allow it in the agent's Access details."
|
||||
),
|
||||
"error_type": "tool_not_allowed",
|
||||
"thought_signature": getattr(call, "thought_signature", None),
|
||||
}
|
||||
|
||||
if require_approval:
|
||||
if self.headless:
|
||||
tool_row_id = str(tool_data.get("id") or tool_id)
|
||||
@@ -982,6 +1276,12 @@ class ToolExecutor:
|
||||
"pause_type": "awaiting_approval",
|
||||
"thought_signature": getattr(call, "thought_signature", None),
|
||||
}
|
||||
# The card shows what will be sent: fixed values replace what the
|
||||
# model asked for. ``arguments`` stays as the model sent it, since
|
||||
# resuming replays it to the model, which never sees fixed values.
|
||||
sent = sent_arguments(action_data, arguments, self._connection_parameters(tool_data))
|
||||
if action_data and sent != arguments:
|
||||
payload["sent_arguments"] = sent
|
||||
# Surface the device id so the approval UI can offer a
|
||||
# "don't ask again" sticky-pattern action for remote devices.
|
||||
if tool_data.get("name") == "remote_device":
|
||||
@@ -992,6 +1292,35 @@ class ToolExecutor:
|
||||
|
||||
return None
|
||||
|
||||
def _on_owner_credentials(self, tool_data: Dict, resolved, action_name: Optional[str] = None) -> bool:
|
||||
"""Whether a call would act with credentials the caller doesn't hold.
|
||||
|
||||
The connection's account when there is one, else the tool owner's
|
||||
stored credentials (see ``holds_owner_credentials``). An API-key
|
||||
caller and any scheduled run hold none of their own: the run acts as
|
||||
the owner. A public-link user in the app holds their own account only.
|
||||
|
||||
Args:
|
||||
tool_data: The ``user_tools`` row being called.
|
||||
resolved: The connection ``resolve_connection`` picked, or None.
|
||||
action_name: The action called; an API tool action carries its
|
||||
own headers and query values.
|
||||
|
||||
Returns:
|
||||
True when the call would use someone else's credentials.
|
||||
"""
|
||||
from docsgpt.connectors.permissions import holds_owner_credentials
|
||||
|
||||
if resolved is not None:
|
||||
holder = (resolved.row or {}).get("user_id")
|
||||
elif holds_owner_credentials(tool_data, action_name):
|
||||
holder = tool_data.get("user_id")
|
||||
else:
|
||||
return False
|
||||
if self.external_caller or self.headless:
|
||||
return True
|
||||
return holder != self.user
|
||||
|
||||
def _remote_device_requires_approval(
|
||||
self,
|
||||
tool_data: Dict,
|
||||
@@ -1302,6 +1631,12 @@ class ToolExecutor:
|
||||
"arguments": call_args,
|
||||
}
|
||||
tool_data = tools_dict[tool_id]
|
||||
# Name the service a connection-backed tool used, so the chip can show
|
||||
# its logo ("Searched Notion"). Never the account behind it.
|
||||
resolved = self._resolve_connection(tool_data)
|
||||
if resolved is not None and resolved.connector_key:
|
||||
tool_call_data["connector_key"] = resolved.connector_key
|
||||
tool_call_data["connector_name"] = resolved.connector_name
|
||||
# Surface the device id on remote_device tool-call events so the
|
||||
# approval UI can wire up the sticky "don't ask again" button.
|
||||
if tool_data.get("name") == "remote_device":
|
||||
@@ -1347,35 +1682,40 @@ class ToolExecutor:
|
||||
else next(action for action in tool_data["actions"] if action["name"] == action_name)
|
||||
)
|
||||
|
||||
query_params, headers, body, parameters = {}, {}, {}, {}
|
||||
param_types = {
|
||||
"query_params": query_params,
|
||||
"headers": headers,
|
||||
"body": body,
|
||||
"parameters": parameters,
|
||||
}
|
||||
if "connector_key" in tool_call_data:
|
||||
from docsgpt.connectors.permissions import action_access
|
||||
|
||||
for param_type, target_dict in param_types.items():
|
||||
if param_type in action_data and action_data[param_type].get("properties"):
|
||||
for param, details in action_data[param_type]["properties"].items():
|
||||
if param not in call_args and "value" in details and details["value"]:
|
||||
target_dict[param] = details["value"]
|
||||
for param, value in call_args.items():
|
||||
for param_type, target_dict in param_types.items():
|
||||
if param_type in action_data and param in action_data[param_type].get("properties", {}):
|
||||
target_dict[param] = value
|
||||
tool_call_data["access"] = action_access(tool_data.get("name"), action_data)
|
||||
|
||||
# Fixed values win over whatever the model sent for the same key.
|
||||
sections = resolve_arguments(action_data, call_args, self._connection_parameters(tool_data))
|
||||
query_params, headers = sections["query_params"], sections["headers"]
|
||||
body, parameters = sections["body"], sections["parameters"]
|
||||
# The chat shows what was sent; ``arguments`` keeps what the model
|
||||
# asked for, which is what a later turn replays to it.
|
||||
sent = sent_arguments(action_data, call_args, self._connection_parameters(tool_data))
|
||||
if sent != (call_args if isinstance(call_args, dict) else {}):
|
||||
tool_call_data["sent_arguments"] = sent
|
||||
|
||||
# Load tool (with caching)
|
||||
tool = self._get_or_load_tool(
|
||||
tool_data,
|
||||
tool_id,
|
||||
action_name,
|
||||
headers=headers,
|
||||
query_params=query_params,
|
||||
)
|
||||
from docsgpt.connectors.service import ConnectionUnavailable
|
||||
|
||||
connection_error = None
|
||||
try:
|
||||
tool = self._get_or_load_tool(
|
||||
tool_data,
|
||||
tool_id,
|
||||
action_name,
|
||||
headers=headers,
|
||||
query_params=query_params,
|
||||
)
|
||||
except ConnectionUnavailable as exc:
|
||||
tool, connection_error = None, str(exc)
|
||||
|
||||
if tool is None:
|
||||
error_message = (
|
||||
error_message = connection_error and (
|
||||
f"{connection_error}. Ask the user to connect it in Settings > Connectors, then try again."
|
||||
) or (
|
||||
f"Failed to load tool '{tool_data.get('name')}' (tool_id key={tool_id}): missing 'id' on tool row."
|
||||
)
|
||||
logger.error(
|
||||
@@ -1535,6 +1875,7 @@ class ToolExecutor:
|
||||
return cached
|
||||
|
||||
tm = ToolManager(config={})
|
||||
load_user = self.user
|
||||
|
||||
if tool_data["name"] == "api_tool":
|
||||
action_config = tool_data["config"]["actions"][action_name]
|
||||
@@ -1549,6 +1890,9 @@ class ToolExecutor:
|
||||
tool_config["body_encoding_rules"] = action_config.get("body_encoding_rules", {})
|
||||
else:
|
||||
tool_config = tool_data["config"].copy() if tool_data["config"] else {}
|
||||
# Whose MCP tokens a tool uses is decided by resolving its
|
||||
# connection below, never by a value stored in its config.
|
||||
tool_config.pop("connection_id", None)
|
||||
# Credentials are PBKDF2-bound to the tool OWNER's sub, not the
|
||||
# invoker's. Decrypt with the tool row's user_id so a team member
|
||||
# running an owner's shared tool authenticates with the owner's
|
||||
@@ -1557,7 +1901,17 @@ class ToolExecutor:
|
||||
# silently decrypt-failing. Falls back to self.user for the
|
||||
# agentless path where the tool row carries no user_id.
|
||||
tool_owner = tool_data.get("user_id") or self.user
|
||||
if tool_config.get("encrypted_credentials") and tool_owner:
|
||||
resolved = self._resolve_connection(tool_data)
|
||||
if tool_data["name"] == "mcp_tool":
|
||||
# MCP OAuth tokens are read as the account the call runs on:
|
||||
# the resolved connection's owner (a member's own in member
|
||||
# mode), else the tool owner, so a shared server without a
|
||||
# connection runs on the owner's sign-in like every other
|
||||
# credential.
|
||||
load_user = ((resolved.row or {}).get("user_id") if resolved else None) or tool_owner
|
||||
if resolved is not None:
|
||||
self._apply_connection(tool_data, tool_id, tool_config, resolved)
|
||||
elif tool_config.get("encrypted_credentials") and tool_owner:
|
||||
if tool_owner != self.user:
|
||||
# Credential delegation: the invoker is running a shared
|
||||
# tool with the owner's secrets. Audit it (the agent-run
|
||||
@@ -1606,12 +1960,12 @@ class ToolExecutor:
|
||||
# falls back to ``origin_conversation_id`` as the schedule's
|
||||
# conversation home.
|
||||
tool_config["agent_id"] = str(self.agent_id) if self.agent_id else None
|
||||
if self.external_caller:
|
||||
# Its runs act as the owner for an API-key caller.
|
||||
tool_config["created_via"] = "api"
|
||||
if tool_data["name"] == "mcp_tool":
|
||||
tool_config["query_mode"] = True
|
||||
|
||||
# MCP OAuth tokens are looked up by user id: a shared server runs on
|
||||
# the owner's connection, like every other credential.
|
||||
load_user = (tool_data.get("user_id") or self.user) if tool_data["name"] == "mcp_tool" else self.user
|
||||
tool = tm.load_tool(
|
||||
tool_data["name"],
|
||||
tool_config=tool_config,
|
||||
@@ -1624,10 +1978,83 @@ class ToolExecutor:
|
||||
|
||||
return tool
|
||||
|
||||
def _apply_connection(self, tool_data: Dict, tool_id: str, tool_config: Dict, resolved) -> None:
|
||||
"""Merge a connection's credentials into ``tool_config``.
|
||||
|
||||
A connection only serves the tools its connector provides (a Telegram
|
||||
bot token never reaches an ntfy tool), and an MCP connection's secret
|
||||
only goes to the server it was stored for.
|
||||
|
||||
Args:
|
||||
tool_data: The ``user_tools`` row being run.
|
||||
tool_id: The tool's key in this run.
|
||||
tool_config: The config the tool is loaded with, updated in place.
|
||||
resolved: The connection :func:`resolve_connection` picked.
|
||||
|
||||
Raises:
|
||||
ConnectionUnavailable: The connection needs signing in again, or
|
||||
does not belong to this tool's connector or server.
|
||||
"""
|
||||
from docsgpt.connectors import catalog, service
|
||||
from docsgpt.connectors.resolve import audit_delegation
|
||||
|
||||
unavailable = service.ConnectionUnavailable(
|
||||
f"{resolved.connector_name or 'This service'} needs to be connected",
|
||||
connection_id=resolved.connection_id,
|
||||
)
|
||||
if not resolved.available or resolved.row is None:
|
||||
raise unavailable
|
||||
tool_name = tool_data.get("name")
|
||||
definition = catalog.get_definition(resolved.connector_key)
|
||||
if definition is None or tool_name not in definition.tool_templates:
|
||||
logger.warning(
|
||||
"tool %s (%s) points at a %s connection", tool_data.get("id") or tool_id, tool_name,
|
||||
resolved.connector_key,
|
||||
)
|
||||
raise unavailable
|
||||
if tool_name == "mcp_tool":
|
||||
# A connection's secret only goes to the server it was stored for:
|
||||
# a custom server's own URL (legacy rows name it in ``provider``),
|
||||
# or a preset's or built-in connector's MCP server (GitHub's token
|
||||
# only ever goes to GitHub's). No server known, nothing is sent.
|
||||
stored_for = catalog.base_url(resolved.row.get("server_url"))
|
||||
provider = str(resolved.row.get("provider") or "")
|
||||
if not stored_for and provider.startswith("mcp:"):
|
||||
stored_for = catalog.base_url(provider[len("mcp:"):])
|
||||
if not stored_for:
|
||||
stored_for = definition.mcp_base_url or ""
|
||||
if not stored_for or stored_for != catalog.base_url(tool_config.get("server_url")):
|
||||
raise unavailable
|
||||
if service.builtin_mcp_config(definition) is not None:
|
||||
# Only the connector's own endpoints: GitHub's write one while
|
||||
# the tool opted in and an admin allows it, else read-only.
|
||||
tool_config["server_url"] = service.builtin_mcp_url(
|
||||
definition, tool_config.get("server_url"), resolved.writes_allowed,
|
||||
)
|
||||
audit_delegation(
|
||||
resolved,
|
||||
invoker=self.user,
|
||||
resource_type="tool",
|
||||
resource_id=str(tool_data.get("id") or tool_id),
|
||||
agent_id=self.agent_id,
|
||||
)
|
||||
tool_config.pop("encrypted_credentials", None)
|
||||
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":
|
||||
# MCP OAuth tokens are read by connection id inside the tool.
|
||||
tool_config["connection_id"] = resolved.connection_id
|
||||
|
||||
# Keys the client needs that are not part of the fixed shape below. They are
|
||||
# small and optional, and are copied only when present so an ordinary tool
|
||||
# call does not grow null columns in every persisted row.
|
||||
_PRESERVED_TOOL_CALL_KEYS = ("artifacts", "device_id")
|
||||
_PRESERVED_TOOL_CALL_KEYS = (
|
||||
"artifacts", "device_id", "connector_key", "connector_name", "access", "sent_arguments",
|
||||
)
|
||||
|
||||
def get_truncated_tool_calls(self) -> List[Dict]:
|
||||
"""Project tool calls into the shape that is streamed and persisted.
|
||||
|
||||
@@ -0,0 +1,334 @@
|
||||
"""Fixed ("pinned") tool parameters.
|
||||
|
||||
Every parameter of a stored tool action carries ``filled_by_llm`` and
|
||||
``value``. A parameter with ``filled_by_llm`` false is hidden from the model;
|
||||
when it also has a value, that value is *pinned*: it is sent on every call and
|
||||
the model can neither see nor override it. An empty string means "no value":
|
||||
the parameter is left out of the call (optional OpenAPI parameters are stored
|
||||
that way), while ``0`` and ``false`` are real values.
|
||||
|
||||
The run-time rules live in :func:`resolve_arguments`; the helpers below
|
||||
validate and apply pin changes coming from the API.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import copy
|
||||
from typing import Any, Iterable, Mapping, Optional
|
||||
|
||||
#: The places an action keeps parameter schemas: ``parameters`` for ordinary
|
||||
#: tools and MCP, the other three for ``api_tool`` requests.
|
||||
PARAM_SECTIONS = ("query_params", "headers", "body", "parameters")
|
||||
|
||||
_NUMERIC = {"integer": int, "number": float}
|
||||
|
||||
|
||||
def _properties(action: Mapping, section: str) -> dict:
|
||||
block = action.get(section)
|
||||
if not isinstance(block, dict):
|
||||
return {}
|
||||
properties = block.get("properties")
|
||||
return properties if isinstance(properties, dict) else {}
|
||||
|
||||
|
||||
def iter_parameters(action: Mapping) -> Iterable[tuple[str, str, dict]]:
|
||||
"""``(section, name, details)`` for every parameter schema of ``action``."""
|
||||
for section in PARAM_SECTIONS:
|
||||
for name, details in _properties(action, section).items():
|
||||
if isinstance(details, dict):
|
||||
yield section, name, details
|
||||
|
||||
|
||||
def llm_fills(details: Mapping) -> bool:
|
||||
"""Whether the model fills this parameter (the default when unset)."""
|
||||
return bool(details.get("filled_by_llm", True))
|
||||
|
||||
|
||||
def has_value(details: Mapping) -> bool:
|
||||
"""Whether a parameter carries a value: anything but a missing, None or empty string."""
|
||||
value = details.get("value")
|
||||
return value is not None and value != ""
|
||||
|
||||
|
||||
def is_pinned(details: Mapping) -> bool:
|
||||
"""Whether the parameter is hidden from the model and always sent with its value."""
|
||||
return not llm_fills(details) and has_value(details)
|
||||
|
||||
|
||||
def pinned_names(action: Mapping) -> list[str]:
|
||||
"""Names of the parameters of ``action`` that carry a pinned value."""
|
||||
return [name for _section, name, details in iter_parameters(action) if is_pinned(details)]
|
||||
|
||||
|
||||
def resolve_arguments(
|
||||
action: Mapping,
|
||||
llm_arguments: Mapping[str, Any],
|
||||
connection_pins: Optional[Mapping[str, Any]] = None,
|
||||
) -> dict[str, dict]:
|
||||
"""The arguments a call runs with, per parameter section.
|
||||
|
||||
Only parameters the action declares are passed. For each one:
|
||||
|
||||
* a value the connection fixes (``connection_pins``, e.g. Telegram's
|
||||
default chat) always wins;
|
||||
* a parameter hidden from the model takes its stored value, or is left
|
||||
out when it has none; whatever the model sent for it is ignored, so a
|
||||
model that names the key anyway (a mistake, or a prompt injection)
|
||||
changes nothing;
|
||||
* any other parameter takes the model's value, falling back to its
|
||||
stored default.
|
||||
|
||||
Args:
|
||||
action: The stored action (``parameters`` / ``query_params`` /
|
||||
``headers`` / ``body`` schemas with ``filled_by_llm`` and ``value``).
|
||||
llm_arguments: The arguments the model sent.
|
||||
connection_pins: Parameter name to a value the connection fixes.
|
||||
|
||||
Returns:
|
||||
Section name to ``{parameter: value}`` for every section in
|
||||
:data:`PARAM_SECTIONS` (empty dicts included).
|
||||
"""
|
||||
connection_pins = connection_pins or {}
|
||||
resolved: dict[str, dict] = {section: {} for section in PARAM_SECTIONS}
|
||||
for section, name, details in iter_parameters(action):
|
||||
target = resolved[section]
|
||||
if name in connection_pins:
|
||||
target[name] = connection_pins[name]
|
||||
elif not llm_fills(details):
|
||||
if has_value(details):
|
||||
target[name] = details["value"]
|
||||
elif name in llm_arguments:
|
||||
target[name] = llm_arguments[name]
|
||||
elif has_value(details):
|
||||
target[name] = details["value"]
|
||||
return resolved
|
||||
|
||||
|
||||
#: What the chat shows in place of a value the owner fixed.
|
||||
FIXED_MASK = "(fixed)"
|
||||
|
||||
|
||||
def sent_arguments(
|
||||
action: Mapping,
|
||||
llm_arguments: Mapping[str, Any],
|
||||
connection_pins: Optional[Mapping[str, Any]] = None,
|
||||
) -> dict:
|
||||
"""What a call sends, flattened for the chat (the approval card, a finished call).
|
||||
|
||||
It follows :func:`resolve_arguments`, but any value that comes from the
|
||||
stored action rather than from the model shows as :data:`FIXED_MASK`:
|
||||
a fixed value, and a default the model left alone, may be a secret (an
|
||||
api_tool query value is decrypted into the action before the call), and
|
||||
the chat is shown to whoever runs the agent, over the API or a widget
|
||||
too. A value the connection sets (Telegram's default chat) is the
|
||||
account's own setting and shows as it is. Headers are left out.
|
||||
|
||||
Args:
|
||||
action: The stored action, with any secrets merged back.
|
||||
llm_arguments: The arguments the model sent.
|
||||
connection_pins: Parameter name to a value the connection fixes.
|
||||
|
||||
Returns:
|
||||
Parameter name to the value shown for it.
|
||||
"""
|
||||
connection_pins = connection_pins or {}
|
||||
shown: dict = {}
|
||||
for section, name, details in iter_parameters(action):
|
||||
if section == "headers":
|
||||
continue
|
||||
if name in connection_pins:
|
||||
shown[name] = connection_pins[name]
|
||||
elif not llm_fills(details):
|
||||
if has_value(details):
|
||||
shown[name] = FIXED_MASK
|
||||
elif name in llm_arguments:
|
||||
shown[name] = llm_arguments[name]
|
||||
elif has_value(details):
|
||||
shown[name] = FIXED_MASK
|
||||
return shown
|
||||
|
||||
|
||||
def coerce_value(details: Mapping, value: Any) -> Any:
|
||||
"""Check a value for a pin against the parameter's declared type.
|
||||
|
||||
Strings typed into a form are converted for ``integer``, ``number`` and
|
||||
``boolean`` parameters.
|
||||
|
||||
Raises:
|
||||
ValueError: The value does not fit the type, or is empty.
|
||||
"""
|
||||
if value is None or value == "":
|
||||
raise ValueError("A fixed value cannot be empty")
|
||||
kind = details.get("type")
|
||||
if isinstance(kind, list):
|
||||
kind = next((k for k in kind if k != "null"), None)
|
||||
if kind in _NUMERIC:
|
||||
if isinstance(value, bool):
|
||||
raise ValueError(f"Expected a {kind}")
|
||||
try:
|
||||
number = _NUMERIC[kind](value)
|
||||
except (TypeError, ValueError):
|
||||
raise ValueError(f"Expected a {kind}") from None
|
||||
if kind == "integer" and isinstance(value, float) and not value.is_integer():
|
||||
raise ValueError("Expected an integer")
|
||||
return number
|
||||
if kind == "boolean":
|
||||
if isinstance(value, bool):
|
||||
return value
|
||||
if isinstance(value, str) and value.strip().lower() in ("true", "false"):
|
||||
return value.strip().lower() == "true"
|
||||
raise ValueError("Expected true or false")
|
||||
if kind == "array":
|
||||
if not isinstance(value, list):
|
||||
raise ValueError("Expected a list")
|
||||
return value
|
||||
if kind == "object":
|
||||
if not isinstance(value, dict):
|
||||
raise ValueError("Expected an object")
|
||||
return value
|
||||
if kind in (None, "string"):
|
||||
if not isinstance(value, (str, int, float, bool)):
|
||||
raise ValueError("Expected text")
|
||||
return str(value) if kind == "string" else value
|
||||
return value
|
||||
|
||||
|
||||
def set_pins(action: Mapping, pins: Mapping[str, Any]) -> dict:
|
||||
"""Return ``action`` with ``pins`` applied.
|
||||
|
||||
Args:
|
||||
action: A stored action.
|
||||
pins: Parameter name to the value to always use, or None to let the
|
||||
model decide again.
|
||||
|
||||
Raises:
|
||||
ValueError: A parameter does not exist on the action, or a value does
|
||||
not fit its type.
|
||||
"""
|
||||
updated = copy.deepcopy(dict(action))
|
||||
known = {name: details for _section, name, details in iter_parameters(updated)}
|
||||
for name, value in pins.items():
|
||||
details = known.get(name)
|
||||
if details is None:
|
||||
raise ValueError(f"Unknown parameter: {name}")
|
||||
if value is None:
|
||||
details["filled_by_llm"] = True
|
||||
details["value"] = ""
|
||||
else:
|
||||
details["value"] = coerce_value(details, value)
|
||||
details["filled_by_llm"] = False
|
||||
return updated
|
||||
|
||||
|
||||
class PinChangeRefused(PermissionError):
|
||||
"""Someone other than the tool's owner tried to change a fixed value."""
|
||||
|
||||
|
||||
_ACTION_KEYS = ("active", "require_approval", "description")
|
||||
_PARAM_KEYS = ("filled_by_llm", "value", "description")
|
||||
|
||||
|
||||
def _pin_state(details: Mapping) -> tuple:
|
||||
return (llm_fills(details), details.get("value") if has_value(details) else None)
|
||||
|
||||
|
||||
def merge_submitted_actions(stored: list, submitted: Any, *, may_change_pins: bool) -> list:
|
||||
"""Validate a full list of actions sent by a client against the stored ones.
|
||||
|
||||
The stored actions are the schema: a client can switch actions on and off,
|
||||
change approval and descriptions and (the owner only) fixed values, but
|
||||
cannot add actions or parameters, or change a parameter's type.
|
||||
|
||||
Args:
|
||||
stored: The tool's stored actions.
|
||||
submitted: What the client sent (untrusted).
|
||||
may_change_pins: False for a team editor: any change to
|
||||
``filled_by_llm`` or ``value`` is refused.
|
||||
|
||||
Returns:
|
||||
The actions to store: every stored action, updated from the
|
||||
submission.
|
||||
|
||||
Raises:
|
||||
ValueError: The submission is malformed or names an unknown action or
|
||||
parameter, or a fixed value does not fit its parameter.
|
||||
PinChangeRefused: ``may_change_pins`` is False and a fixed value changed.
|
||||
"""
|
||||
if not isinstance(submitted, list):
|
||||
raise ValueError("actions must be a list")
|
||||
by_name = {a.get("name"): a for a in stored if isinstance(a, dict) and a.get("name")}
|
||||
seen: set = set()
|
||||
updates: dict[str, dict] = {}
|
||||
for entry in submitted:
|
||||
if not isinstance(entry, dict) or not isinstance(entry.get("name"), str):
|
||||
raise ValueError("Each action needs a name")
|
||||
name = entry["name"]
|
||||
if name not in by_name:
|
||||
raise ValueError(f"Unknown action: {name}")
|
||||
if name in seen:
|
||||
raise ValueError(f"Duplicate action: {name}")
|
||||
seen.add(name)
|
||||
updates[name] = entry
|
||||
|
||||
merged = []
|
||||
for action in stored:
|
||||
if not isinstance(action, dict) or action.get("name") not in updates:
|
||||
merged.append(action)
|
||||
continue
|
||||
entry = updates[action["name"]]
|
||||
result = copy.deepcopy(action)
|
||||
for key in _ACTION_KEYS:
|
||||
if key in entry:
|
||||
value = entry[key]
|
||||
if key == "description" and not isinstance(value, str):
|
||||
raise ValueError("description must be text")
|
||||
result[key] = bool(value) if key != "description" else value
|
||||
for section in PARAM_SECTIONS:
|
||||
sent = _properties(entry, section)
|
||||
if not sent:
|
||||
continue
|
||||
properties = _properties(result, section)
|
||||
for param, sent_details in sent.items():
|
||||
details = properties.get(param)
|
||||
if details is None or not isinstance(details, dict):
|
||||
raise ValueError(f"Unknown parameter {param} on {action['name']}")
|
||||
if not isinstance(sent_details, dict):
|
||||
raise ValueError(f"Parameter {param} must be an object")
|
||||
before = _pin_state(details)
|
||||
for key in _PARAM_KEYS:
|
||||
if key in sent_details:
|
||||
details[key] = sent_details[key]
|
||||
details["filled_by_llm"] = llm_fills(details)
|
||||
if _pin_state(details) == before:
|
||||
continue
|
||||
if not may_change_pins:
|
||||
raise PinChangeRefused("Only the tool's owner can change fixed values")
|
||||
if is_pinned(details):
|
||||
details["value"] = coerce_value(details, details["value"])
|
||||
merged.append(result)
|
||||
return merged
|
||||
|
||||
|
||||
def carry_pins_between(previous: Optional[list], fresh: list) -> list:
|
||||
""":func:`carry_pins` for whole action lists, matching actions by name."""
|
||||
old = {a.get("name"): a for a in previous or [] if isinstance(a, dict)}
|
||||
return [carry_pins(old.get(a.get("name")), a) if isinstance(a, dict) else a for a in fresh]
|
||||
|
||||
|
||||
def carry_pins(previous: Optional[Mapping], fresh: Mapping) -> dict:
|
||||
"""Copy fixed values from an action's old copy onto its re-discovered one.
|
||||
|
||||
Parameters that still exist keep ``filled_by_llm`` and ``value``; removed
|
||||
ones are gone with the new schema.
|
||||
"""
|
||||
result = copy.deepcopy(dict(fresh))
|
||||
if not previous:
|
||||
return result
|
||||
old = {(section, name): details for section, name, details in iter_parameters(previous)}
|
||||
for section, name, details in iter_parameters(result):
|
||||
before = old.get((section, name))
|
||||
if before is None:
|
||||
continue
|
||||
details["filled_by_llm"] = llm_fills(before)
|
||||
details["value"] = before.get("value", "")
|
||||
return result
|
||||
@@ -136,6 +136,7 @@ class BraveSearchTool(Tool):
|
||||
return [
|
||||
{
|
||||
"name": "brave_web_search",
|
||||
"access": "read",
|
||||
"description": (
|
||||
"Search the web with Brave Search. Returns result titles, "
|
||||
"URLs, and snippets. Use it for current events or "
|
||||
@@ -163,6 +164,7 @@ class BraveSearchTool(Tool):
|
||||
},
|
||||
{
|
||||
"name": "brave_image_search",
|
||||
"access": "read",
|
||||
"description": (
|
||||
"Search for images with Brave Search. Returns image "
|
||||
"titles, page URLs, and thumbnail URLs."
|
||||
|
||||
@@ -1,10 +1,11 @@
|
||||
import asyncio
|
||||
import base64
|
||||
import concurrent.futures
|
||||
import hashlib
|
||||
import json
|
||||
import logging
|
||||
import time
|
||||
from typing import Any, Dict, List, Optional
|
||||
from typing import Any, Callable, Dict, List, Optional
|
||||
from urllib.parse import parse_qs, urlparse
|
||||
|
||||
from fastmcp import Client
|
||||
@@ -15,7 +16,12 @@ from fastmcp.client.transports import (
|
||||
StreamableHttpTransport,
|
||||
)
|
||||
from mcp.client.auth import OAuthClientProvider, TokenStorage
|
||||
from mcp.shared.auth import OAuthClientInformationFull, OAuthClientMetadata, OAuthToken
|
||||
from mcp.shared.auth import (
|
||||
AuthorizationCodeResult,
|
||||
OAuthClientInformationFull,
|
||||
OAuthClientMetadata,
|
||||
OAuthToken,
|
||||
)
|
||||
from pydantic import AnyHttpUrl, ValidationError
|
||||
from redis import Redis
|
||||
|
||||
@@ -32,12 +38,54 @@ logger = logging.getLogger(__name__)
|
||||
_mcp_clients_cache = {}
|
||||
|
||||
|
||||
def forget_cached_clients(*identities: str) -> None:
|
||||
"""Drop cached OAuth clients signed in as any of ``identities``.
|
||||
|
||||
OAuth cache keys name the connection or user whose tokens the client
|
||||
holds (see ``MCPTool._generate_cache_key``).
|
||||
|
||||
Args:
|
||||
identities: Connection ids and user ids.
|
||||
"""
|
||||
markers = tuple(f"#oauth:{identity}:" for identity in identities if identity)
|
||||
for key in [k for k in list(_mcp_clients_cache) if any(marker in k for marker in markers)]:
|
||||
_mcp_clients_cache.pop(key, None)
|
||||
|
||||
# A token expiry long past: a stored token of unknown age is renewed before use.
|
||||
_EXPIRED = 1.0
|
||||
|
||||
_ANNOTATION_HINTS = ("readOnlyHint", "destructiveHint", "idempotentHint", "openWorldHint")
|
||||
|
||||
|
||||
def _secret_fingerprint(secret: str) -> str:
|
||||
"""A short digest of a whole secret, for telling cached clients apart.
|
||||
|
||||
A prefix of the secret itself is not enough: every fine-grained GitHub
|
||||
token starts with ``github_pat_``, so two users' tokens would share (and
|
||||
reuse) one cached client carrying the first user's token.
|
||||
"""
|
||||
return hashlib.sha256(secret.encode("utf-8")).hexdigest()[:16]
|
||||
|
||||
|
||||
def _annotation_hints(annotations: Any) -> Dict[str, bool]:
|
||||
"""The boolean MCP tool annotation hints, from a model or a dict."""
|
||||
if annotations is None:
|
||||
return {}
|
||||
if hasattr(annotations, "model_dump"):
|
||||
annotations = annotations.model_dump()
|
||||
if not isinstance(annotations, dict):
|
||||
return {}
|
||||
return {k: annotations[k] for k in _ANNOTATION_HINTS if isinstance(annotations.get(k), bool)}
|
||||
|
||||
|
||||
class MCPTool(Tool):
|
||||
"""
|
||||
MCP Tool
|
||||
Connect to remote Model Context Protocol (MCP) servers to access dynamic tools and resources.
|
||||
"""
|
||||
|
||||
connection_id: Optional[str] = None
|
||||
|
||||
def __init__(self, config: Dict[str, Any], user_id: Optional[str] = None):
|
||||
"""
|
||||
Initialize the MCP Tool with configuration.
|
||||
@@ -76,6 +124,10 @@ class MCPTool(Tool):
|
||||
self.oauth_scopes = config.get("oauth_scopes", [])
|
||||
self.oauth_task_id = config.get("oauth_task_id", None)
|
||||
self.oauth_client_name = config.get("oauth_client_name", "DocsGPT-MCP")
|
||||
# The connection whose OAuth tokens this tool uses. Set by the tool
|
||||
# executor for connection-backed tools, so a shared tool in ``owner``
|
||||
# mode signs in with its owner's account rather than the invoker's.
|
||||
self.connection_id = config.get("connection_id")
|
||||
self.redirect_uri = self._resolve_redirect_uri(config.get("redirect_uri"))
|
||||
# Pulled out of ``config`` (rather than left in ``self.config``)
|
||||
# because it is a callable supplied by the OAuth worker — not
|
||||
@@ -105,13 +157,16 @@ class MCPTool(Tool):
|
||||
raise ValueError(f"Invalid MCP server URL: {exc}") from exc
|
||||
|
||||
def _resolve_redirect_uri(self, configured_redirect_uri: Optional[str]) -> str:
|
||||
if configured_redirect_uri:
|
||||
return configured_redirect_uri.rstrip("/")
|
||||
|
||||
# The operator's setting wins over the page's own origin: a page opened
|
||||
# on a plain-HTTP address would otherwise register a callback that
|
||||
# servers such as Linear refuse.
|
||||
explicit = settings.MCP_OAUTH_REDIRECT_URI
|
||||
if explicit:
|
||||
return explicit.rstrip("/")
|
||||
|
||||
if configured_redirect_uri:
|
||||
return configured_redirect_uri.rstrip("/")
|
||||
|
||||
connector_base = settings.CONNECTOR_REDIRECT_BASE_URI
|
||||
if connector_base:
|
||||
parsed = urlparse(connector_base)
|
||||
@@ -125,7 +180,9 @@ class MCPTool(Tool):
|
||||
auth_key = ""
|
||||
if self.auth_type == "oauth":
|
||||
scopes_str = ",".join(self.oauth_scopes) if self.oauth_scopes else "none"
|
||||
oauth_identity = self.user_id or self.oauth_task_id or "anonymous"
|
||||
# A connection-backed tool shares a client only with calls that
|
||||
# use the same connection's tokens.
|
||||
oauth_identity = self.connection_id or self.user_id or self.oauth_task_id or "anonymous"
|
||||
auth_key = (
|
||||
f"oauth:{oauth_identity}:{self.oauth_client_name}:{scopes_str}:{self.redirect_uri}"
|
||||
)
|
||||
@@ -133,10 +190,10 @@ class MCPTool(Tool):
|
||||
token = self.auth_credentials.get(
|
||||
"bearer_token", ""
|
||||
) or self.auth_credentials.get("access_token", "")
|
||||
auth_key = f"bearer:{token[:10]}..." if token else "bearer:none"
|
||||
auth_key = f"bearer:{_secret_fingerprint(token)}" if token else "bearer:none"
|
||||
elif self.auth_type == "api_key":
|
||||
api_key = self.auth_credentials.get("api_key", "")
|
||||
auth_key = f"apikey:{api_key[:10]}..." if api_key else "apikey:none"
|
||||
auth_key = f"apikey:{_secret_fingerprint(api_key)}" if api_key else "apikey:none"
|
||||
elif self.auth_type == "basic":
|
||||
username = self.auth_credentials.get("username", "")
|
||||
auth_key = f"basic:{username}"
|
||||
@@ -165,6 +222,7 @@ class MCPTool(Tool):
|
||||
redis_client=redis_client,
|
||||
redirect_uri=self.redirect_uri,
|
||||
user_id=self.user_id,
|
||||
connection_id=self.connection_id,
|
||||
)
|
||||
else:
|
||||
auth = DocsGPTOAuth(
|
||||
@@ -175,6 +233,7 @@ class MCPTool(Tool):
|
||||
task_id=self.oauth_task_id,
|
||||
user_id=self.user_id,
|
||||
redirect_publish=self.oauth_redirect_publish,
|
||||
connection_id=self.connection_id,
|
||||
)
|
||||
elif self.auth_type == "bearer":
|
||||
token = self.auth_credentials.get(
|
||||
@@ -245,6 +304,9 @@ class MCPTool(Tool):
|
||||
}
|
||||
if hasattr(tool, "inputSchema"):
|
||||
tool_dict["inputSchema"] = tool.inputSchema
|
||||
annotations = _annotation_hints(getattr(tool, "annotations", None))
|
||||
if annotations:
|
||||
tool_dict["annotations"] = annotations
|
||||
tools_dict.append(tool_dict)
|
||||
elif isinstance(tool, dict):
|
||||
tools_dict.append(tool)
|
||||
@@ -493,7 +555,7 @@ class MCPTool(Tool):
|
||||
|
||||
def _test_oauth_connection(self) -> Dict:
|
||||
storage = DBTokenStorage(
|
||||
server_url=self.server_url, user_id=self.user_id,
|
||||
server_url=self.server_url, user_id=self.user_id, connection_id=self.connection_id,
|
||||
)
|
||||
loop = asyncio.new_event_loop()
|
||||
try:
|
||||
@@ -580,6 +642,11 @@ class MCPTool(Tool):
|
||||
"description": tool.get("description", ""),
|
||||
"parameters": parameters_schema,
|
||||
}
|
||||
# ``readOnlyHint`` / ``destructiveHint`` decide whether the action
|
||||
# is a read (always allowed) or a write (needs approval).
|
||||
annotations = _annotation_hints(tool.get("annotations"))
|
||||
if annotations:
|
||||
action["annotations"] = annotations
|
||||
actions.append(action)
|
||||
return actions
|
||||
|
||||
@@ -670,6 +737,10 @@ class MCPTool(Tool):
|
||||
}
|
||||
|
||||
|
||||
class MCPReauthorizationRequired(Exception):
|
||||
"""An MCP sign-in expired and could not be renewed: its owner must sign in again."""
|
||||
|
||||
|
||||
class DocsGPTOAuth(OAuthClientProvider):
|
||||
"""
|
||||
Custom OAuth handler for DocsGPT that uses frontend redirect instead of browser.
|
||||
@@ -688,6 +759,7 @@ class DocsGPTOAuth(OAuthClientProvider):
|
||||
additional_client_metadata: dict[str, Any] | None = None,
|
||||
skip_redirect_validation: bool = False,
|
||||
redirect_publish=None,
|
||||
connection_id: Optional[str] = None,
|
||||
):
|
||||
self.redirect_uri = redirect_uri
|
||||
self.redis_client = redis_client
|
||||
@@ -717,19 +789,70 @@ class DocsGPTOAuth(OAuthClientProvider):
|
||||
server_url=self.server_base_url,
|
||||
user_id=self.user_id,
|
||||
expected_redirect_uri=None if skip_redirect_validation else redirect_uri,
|
||||
connection_id=connection_id,
|
||||
)
|
||||
|
||||
# The SDK checks the server's protected-resource metadata against a
|
||||
# resource derived from this URL, so it gets the full MCP endpoint:
|
||||
# Linear and Sentry publish ``https://host/mcp``, and an origin-only
|
||||
# URL fails that check before sign-in. Tokens stay keyed by origin.
|
||||
super().__init__(
|
||||
server_url=self.server_base_url,
|
||||
server_url=mcp_url.rstrip("/") or self.server_base_url,
|
||||
client_metadata=client_metadata,
|
||||
storage=storage,
|
||||
redirect_handler=self.redirect_handler,
|
||||
callback_handler=self.callback_handler,
|
||||
)
|
||||
self.context.prepare_token_auth = self._one_client_authentication(self.context.prepare_token_auth)
|
||||
|
||||
self.auth_url = None
|
||||
self.extracted_state = None
|
||||
|
||||
@staticmethod
|
||||
def _one_client_authentication(prepare: Callable) -> Callable:
|
||||
"""Wrap the SDK's token-request auth so the client authenticates one way.
|
||||
|
||||
With ``client_secret_basic`` the SDK sets the Basic header but leaves
|
||||
``client_id`` in the body, which servers such as Linear reject as a
|
||||
second method; RFC 6749 names the client in the body only when it does
|
||||
not authenticate otherwise.
|
||||
|
||||
Args:
|
||||
prepare: The SDK context's ``prepare_token_auth``.
|
||||
|
||||
Returns:
|
||||
The same function, minus ``client_id`` in the body under Basic auth.
|
||||
"""
|
||||
|
||||
def prepare_token_auth(
|
||||
data: dict[str, str], headers: dict[str, str] | None = None
|
||||
) -> tuple[dict[str, str], dict[str, str]]:
|
||||
data, headers = prepare(data, headers)
|
||||
if headers.get("Authorization", "").startswith("Basic "):
|
||||
data = {k: v for k, v in data.items() if k != "client_id"}
|
||||
return data, headers
|
||||
|
||||
return prepare_token_auth
|
||||
|
||||
async def _initialize(self) -> None:
|
||||
"""Load the stored sign-in along with when its access token expires.
|
||||
|
||||
The SDK learns a token's expiry only from a token response in this
|
||||
process, so a token read back from the connection gets the expiry
|
||||
saved with it, and an expired one is renewed with its refresh token
|
||||
before it is sent. One saved before expiries were kept is of unknown
|
||||
age and is renewed once.
|
||||
"""
|
||||
await super()._initialize()
|
||||
tokens = self.context.current_tokens
|
||||
if tokens is None:
|
||||
return
|
||||
expires_at = getattr(self.context.storage, "expires_at", None)
|
||||
if expires_at:
|
||||
self.context.token_expiry_time = float(expires_at)
|
||||
elif tokens.refresh_token and tokens.expires_in:
|
||||
self.context.token_expiry_time = _EXPIRED
|
||||
|
||||
def _process_auth_url(self, authorization_url: str) -> tuple[str, str]:
|
||||
"""Process authorization URL to extract state"""
|
||||
try:
|
||||
@@ -771,8 +894,13 @@ class DocsGPTOAuth(OAuthClientProvider):
|
||||
exc_info=True,
|
||||
)
|
||||
|
||||
async def callback_handler(self) -> tuple[str, str | None]:
|
||||
"""Wait for auth code from Redis using the state value."""
|
||||
async def callback_handler(self) -> AuthorizationCodeResult:
|
||||
"""Wait for auth code from Redis using the state value.
|
||||
|
||||
Returns:
|
||||
The code, the state it came back with, and the RFC 9207 issuer
|
||||
when the authorization server sent one.
|
||||
"""
|
||||
if not self.redis_client or not self.extracted_state:
|
||||
raise Exception("Redis client or state not configured for OAuth")
|
||||
poll_interval = 1
|
||||
@@ -785,15 +913,22 @@ class DocsGPTOAuth(OAuthClientProvider):
|
||||
if code_data:
|
||||
code = code_data.decode()
|
||||
returned_state = self.extracted_state
|
||||
iss_key = f"{self.redis_prefix}iss:{self.extracted_state}"
|
||||
iss_data = self.redis_client.get(iss_key)
|
||||
|
||||
self.redis_client.delete(code_key)
|
||||
self.redis_client.delete(iss_key)
|
||||
self.redis_client.delete(
|
||||
f"{self.redis_prefix}auth_url:{self.extracted_state}"
|
||||
)
|
||||
self.redis_client.delete(
|
||||
f"{self.redis_prefix}state:{self.extracted_state}"
|
||||
)
|
||||
return code, returned_state
|
||||
return AuthorizationCodeResult(
|
||||
code=code,
|
||||
state=returned_state,
|
||||
iss=iss_data.decode() if iss_data else None,
|
||||
)
|
||||
error_key = f"{self.redis_prefix}error:{self.extracted_state}"
|
||||
error_data = self.redis_client.get(error_key)
|
||||
if error_data:
|
||||
@@ -825,26 +960,39 @@ class NonInteractiveOAuth(DocsGPTOAuth):
|
||||
super().__init__(**kwargs)
|
||||
|
||||
async def redirect_handler(self, authorization_url: str) -> None:
|
||||
raise Exception(
|
||||
raise MCPReauthorizationRequired(
|
||||
"OAuth session expired — please re-authorize this MCP server in tool settings."
|
||||
)
|
||||
|
||||
async def callback_handler(self) -> tuple[str, str | None]:
|
||||
raise Exception(
|
||||
async def callback_handler(self) -> AuthorizationCodeResult:
|
||||
raise MCPReauthorizationRequired(
|
||||
"OAuth session expired — please re-authorize this MCP server in tool settings."
|
||||
)
|
||||
|
||||
|
||||
class DBTokenStorage(TokenStorage):
|
||||
"""MCP OAuth tokens and client registration, kept encrypted on the connection.
|
||||
|
||||
Reads and writes go through ``docsgpt.connectors.service``, which stores
|
||||
them in the connection's owner-bound ``encrypted_credentials``. A tool
|
||||
that runs with a specific connection (``owner`` mode on a shared tool)
|
||||
passes its ``connection_id``; otherwise the invoking user's connection
|
||||
for the server's base URL is used.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
server_url: str,
|
||||
user_id: str,
|
||||
expected_redirect_uri: Optional[str] = None,
|
||||
connection_id: Optional[str] = None,
|
||||
):
|
||||
self.server_url = server_url
|
||||
self.user_id = user_id
|
||||
self.expected_redirect_uri = expected_redirect_uri
|
||||
self.connection_id = connection_id
|
||||
# When the stored access token expires (epoch seconds), once read.
|
||||
self.expires_at: Optional[float] = None
|
||||
|
||||
@staticmethod
|
||||
def get_base_url(url: str) -> str:
|
||||
@@ -855,31 +1003,18 @@ class DBTokenStorage(TokenStorage):
|
||||
return f"mcp:{self.get_base_url(self.server_url)}"
|
||||
|
||||
def _fetch_session_data(self) -> dict:
|
||||
"""Read the JSONB ``session_data`` blob for this MCP server row."""
|
||||
from docsgpt.storage.db.repositories.connector_sessions import (
|
||||
ConnectorSessionsRepository,
|
||||
)
|
||||
from docsgpt.storage.db.session import db_readonly
|
||||
"""The decrypted ``tokens`` / ``client_info`` for this MCP server."""
|
||||
from docsgpt.connectors import service
|
||||
|
||||
base_url = self.get_base_url(self.server_url)
|
||||
with db_readonly() as conn:
|
||||
row = ConnectorSessionsRepository(conn).get_by_user_and_server_url(
|
||||
self.user_id, base_url,
|
||||
)
|
||||
if not row:
|
||||
return {}
|
||||
data = row.get("session_data") or {}
|
||||
if isinstance(data, str):
|
||||
try:
|
||||
data = json.loads(data)
|
||||
except ValueError:
|
||||
return {}
|
||||
return data if isinstance(data, dict) else {}
|
||||
return service.read_mcp_secrets(
|
||||
self.user_id, self.get_base_url(self.server_url), self.connection_id,
|
||||
)
|
||||
|
||||
async def get_tokens(self) -> OAuthToken | None:
|
||||
data = await asyncio.to_thread(self._fetch_session_data)
|
||||
if not data or "tokens" not in data:
|
||||
return None
|
||||
self.expires_at = data.get("tokens_expires_at")
|
||||
try:
|
||||
return OAuthToken.model_validate(data["tokens"])
|
||||
except ValidationError as e:
|
||||
@@ -887,22 +1022,19 @@ class DBTokenStorage(TokenStorage):
|
||||
return None
|
||||
|
||||
def _merge(self, patch: dict) -> None:
|
||||
"""Shallow-merge ``patch`` into this row's ``session_data``.
|
||||
"""Merge ``patch`` into the connection's secrets; ``None`` drops a key."""
|
||||
from docsgpt.connectors import service
|
||||
|
||||
Threads ``server_url`` through to the repository so it lands in
|
||||
the scalar column — ``get_by_user_and_server_url`` needs that to
|
||||
resolve the row (``NULL = 'https://...'`` is UNKNOWN in SQL).
|
||||
"""
|
||||
from docsgpt.storage.db.repositories.connector_sessions import (
|
||||
ConnectorSessionsRepository,
|
||||
status = None
|
||||
if patch.get("tokens"):
|
||||
status = service.STATUS_CONNECTED
|
||||
service.update_mcp_secrets(
|
||||
self.user_id,
|
||||
self.get_base_url(self.server_url),
|
||||
patch,
|
||||
connection_id=self.connection_id,
|
||||
status=status,
|
||||
)
|
||||
from docsgpt.storage.db.session import db_session
|
||||
|
||||
base_url = self.get_base_url(self.server_url)
|
||||
with db_session() as conn:
|
||||
ConnectorSessionsRepository(conn).merge_session_data(
|
||||
self.user_id, self._pg_provider(), base_url, patch,
|
||||
)
|
||||
|
||||
def _delete(self) -> None:
|
||||
from docsgpt.storage.db.repositories.connector_sessions import (
|
||||
@@ -918,7 +1050,10 @@ class DBTokenStorage(TokenStorage):
|
||||
async def set_tokens(self, tokens: OAuthToken) -> None:
|
||||
base_url = self.get_base_url(self.server_url)
|
||||
token_dump = tokens.model_dump()
|
||||
await asyncio.to_thread(self._merge, {"tokens": token_dump})
|
||||
# Kept beside the token: a later process cannot tell from
|
||||
# ``expires_in`` alone whether it has expired.
|
||||
self.expires_at = time.time() + tokens.expires_in if tokens.expires_in else None
|
||||
await asyncio.to_thread(self._merge, {"tokens": token_dump, "tokens_expires_at": self.expires_at})
|
||||
logger.info("Saved tokens for %s", base_url)
|
||||
|
||||
async def get_client_info(self) -> OAuthClientInformationFull | None:
|
||||
@@ -1003,7 +1138,7 @@ class MCPOAuthManager:
|
||||
self.redis_prefix = redis_prefix
|
||||
|
||||
def handle_oauth_callback(
|
||||
self, state: str, code: str, error: Optional[str] = None
|
||||
self, state: str, code: str, error: Optional[str] = None, iss: Optional[str] = None
|
||||
) -> bool:
|
||||
"""
|
||||
Handle OAuth callback from provider.
|
||||
@@ -1012,6 +1147,7 @@ class MCPOAuthManager:
|
||||
state: The state parameter from OAuth callback
|
||||
code: The authorization code from OAuth callback
|
||||
error: Error message if OAuth failed
|
||||
iss: The RFC 9207 issuer from the callback, if the server sent one
|
||||
|
||||
Returns:
|
||||
True if successful, False otherwise
|
||||
@@ -1023,6 +1159,9 @@ class MCPOAuthManager:
|
||||
error_key = f"{self.redis_prefix}error:{state}"
|
||||
self.redis_client.setex(error_key, 300, error)
|
||||
raise Exception(f"OAuth error received: {error}")
|
||||
# The issuer goes first: the waiting sign-in reads it once the code lands.
|
||||
if iss:
|
||||
self.redis_client.setex(f"{self.redis_prefix}iss:{state}", 300, iss)
|
||||
code_key = f"{self.redis_prefix}code:{state}"
|
||||
self.redis_client.setex(code_key, 300, code)
|
||||
|
||||
|
||||
@@ -89,6 +89,7 @@ class NtfyTool(Tool):
|
||||
return [
|
||||
{
|
||||
"name": "ntfy_send_message",
|
||||
"access": "write",
|
||||
"description": (
|
||||
"Send a push notification to an ntfy topic on the "
|
||||
"configured server. Provide the message text; title and "
|
||||
|
||||
@@ -136,6 +136,7 @@ class PostgresTool(Tool):
|
||||
return [
|
||||
{
|
||||
"name": "postgres_execute_sql",
|
||||
"access": "write",
|
||||
"description": "Execute an SQL query against the PostgreSQL database and return the results. Use this tool to interact with the database, e.g., retrieve specific data or perform updates. Only SELECT queries will return data, other queries will return execution status.",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
@@ -151,6 +152,7 @@ class PostgresTool(Tool):
|
||||
},
|
||||
{
|
||||
"name": "postgres_get_schema",
|
||||
"access": "read",
|
||||
"description": "Retrieve the schema of the PostgreSQL database, including tables and their columns. Use this to understand the database structure before executing queries. db_name is 'default' if not provided.",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
|
||||
@@ -44,6 +44,9 @@ class SchedulerTool(Tool):
|
||||
self.user_id: Optional[str] = user_id
|
||||
self.agent_id: Optional[str] = cfg.get("agent_id")
|
||||
self.conversation_id: Optional[str] = cfg.get("conversation_id")
|
||||
# ``api`` when an API-key caller asked (set by the executor, never the
|
||||
# model): the schedule's runs then keep that caller's write limits.
|
||||
self.created_via: str = "api" if cfg.get("created_via") == "api" else "chat"
|
||||
|
||||
def execute_action(self, action_name: str, **kwargs: Any) -> str:
|
||||
"""Dispatch on the LLM-supplied action name."""
|
||||
@@ -191,7 +194,7 @@ class SchedulerTool(Tool):
|
||||
timezone=tz or "UTC",
|
||||
tool_allowlist=allowlist,
|
||||
origin_conversation_id=self.conversation_id,
|
||||
created_via="chat",
|
||||
created_via=self.created_via,
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.exception("schedule_task create failed: %s", exc)
|
||||
|
||||
@@ -11,12 +11,22 @@ class TelegramTool(Tool):
|
||||
"""
|
||||
Telegram Bot
|
||||
A flexible Telegram tool for performing various actions (e.g., sending messages, images).
|
||||
Requires a bot token and chat ID for configuration
|
||||
Requires a bot token; a default chat ID set on the connection is used when no chat is named
|
||||
"""
|
||||
|
||||
def __init__(self, config):
|
||||
self.config = config
|
||||
self.token = config.get("token", "")
|
||||
self.default_chat_id = config.get("chat_id") or None
|
||||
|
||||
def _no_chat(self):
|
||||
return {
|
||||
"status": "error",
|
||||
"error": (
|
||||
"No chat to send to. Name a chat_id, or set a default chat ID on the Telegram "
|
||||
"connection in Settings > Connectors."
|
||||
),
|
||||
}
|
||||
|
||||
def execute_action(self, action_name, **kwargs):
|
||||
actions = {
|
||||
@@ -27,14 +37,20 @@ class TelegramTool(Tool):
|
||||
raise ValueError(f"Unknown action: {action_name}")
|
||||
return actions[action_name](**kwargs)
|
||||
|
||||
def _send_message(self, text, chat_id):
|
||||
def _send_message(self, text, chat_id=None):
|
||||
chat_id = chat_id or self.default_chat_id
|
||||
if not chat_id:
|
||||
return self._no_chat()
|
||||
logger.debug("Sending Telegram message to chat_id=%s", chat_id)
|
||||
url = f"https://api.telegram.org/bot{self.token}/sendMessage"
|
||||
payload = {"chat_id": chat_id, "text": text}
|
||||
response = requests.post(url, data=payload, timeout=100)
|
||||
return {"status_code": response.status_code, "message": "Message sent"}
|
||||
|
||||
def _send_image(self, image_url, chat_id):
|
||||
def _send_image(self, image_url, chat_id=None):
|
||||
chat_id = chat_id or self.default_chat_id
|
||||
if not chat_id:
|
||||
return self._no_chat()
|
||||
logger.debug("Sending Telegram image to chat_id=%s", chat_id)
|
||||
url = f"https://api.telegram.org/bot{self.token}/sendPhoto"
|
||||
payload = {"chat_id": chat_id, "photo": image_url}
|
||||
@@ -45,6 +61,7 @@ class TelegramTool(Tool):
|
||||
return [
|
||||
{
|
||||
"name": "telegram_send_message",
|
||||
"access": "write",
|
||||
"description": (
|
||||
"Send a text message to the configured Telegram chat via "
|
||||
"the bot. Compose the final message text before sending."
|
||||
@@ -67,6 +84,7 @@ class TelegramTool(Tool):
|
||||
},
|
||||
{
|
||||
"name": "telegram_send_image",
|
||||
"access": "write",
|
||||
"description": (
|
||||
"Send an image to the configured Telegram chat. Requires "
|
||||
"a publicly accessible image URL."
|
||||
|
||||
@@ -3,6 +3,7 @@ from typing import Any, Dict, List, Optional
|
||||
|
||||
from docsgpt.agents.tools.base import Tool
|
||||
from docsgpt.agents.tools.path_utils import validate_tool_path
|
||||
from docsgpt.storage.db.repositories.sources import SourcesRepository
|
||||
from docsgpt.storage.db.repositories.wiki_pages import (
|
||||
WikiPageConflict,
|
||||
WikiPagesRepository,
|
||||
@@ -21,6 +22,11 @@ MAX_WIKI_PAGE_BYTES = 1_000_000
|
||||
|
||||
_WRITE_ACTIONS = frozenset({"create", "str_replace", "insert", "delete", "rename"})
|
||||
|
||||
OUTSIDE_EDITS_DENIED = (
|
||||
"Error: This wiki's owner doesn't let API or widget users edit it, "
|
||||
"so it can't be changed from here. You can still read it."
|
||||
)
|
||||
|
||||
|
||||
class WikiTool(Tool):
|
||||
"""Wiki
|
||||
@@ -39,6 +45,9 @@ class WikiTool(Tool):
|
||||
self.config = config
|
||||
self.source_id: Optional[str] = config.get("source_id")
|
||||
self.source_owner_id: Optional[str] = config.get("source_owner_id")
|
||||
# An API-key or widget run (it acts as the agent's owner): it writes
|
||||
# only while the wiki's owner allows such edits.
|
||||
self.outside_caller: bool = bool(config.get("outside_caller"))
|
||||
decoded_token = config.get("decoded_token") or {}
|
||||
self.updated_by: Optional[str] = (
|
||||
(decoded_token.get("sub") if decoded_token else None)
|
||||
@@ -228,8 +237,30 @@ class WikiTool(Tool):
|
||||
return message
|
||||
if access is None or not access.can("edit"):
|
||||
return message
|
||||
if self.outside_caller and not self._outside_edits_allowed():
|
||||
return OUTSIDE_EDITS_DENIED
|
||||
return None
|
||||
|
||||
def _outside_edits_allowed(self) -> bool:
|
||||
"""Read the wiki's live ``wiki_outside_edits`` setting.
|
||||
|
||||
Read on every write rather than trusted from the run's setup, so the
|
||||
owner turning it off stops a conversation that is already going.
|
||||
Fails closed on a missing row or a failed lookup.
|
||||
|
||||
Returns:
|
||||
bool: Whether an API-key or widget run may edit the wiki.
|
||||
"""
|
||||
try:
|
||||
with db_readonly() as conn:
|
||||
row = SourcesRepository(conn).get_by_id(str(self.source_id))
|
||||
except Exception:
|
||||
logger.exception(
|
||||
"Wiki outside-edits check failed for source %s", self.source_id
|
||||
)
|
||||
return False
|
||||
return outside_edits_allowed(row)
|
||||
|
||||
def get_config_requirements(self) -> Dict[str, Any]:
|
||||
return {}
|
||||
|
||||
@@ -500,12 +531,21 @@ class WikiTool(Tool):
|
||||
return f"Renamed: {validated_old} -> {validated_new}"
|
||||
|
||||
|
||||
def build_wiki_tool_entry() -> Dict[str, Any]:
|
||||
"""Build the synthetic tools_dict entry for the WikiTool."""
|
||||
def build_wiki_tool_entry(writes_allowed: bool = True, approval_required: bool = False) -> Dict[str, Any]:
|
||||
"""Build the synthetic tools_dict entry for the WikiTool.
|
||||
|
||||
Args:
|
||||
writes_allowed: False offers the model only ``wiki_view``.
|
||||
approval_required: Every write action waits for the caller's approval.
|
||||
"""
|
||||
entry = {"name": "wiki"}
|
||||
entry["actions"] = [
|
||||
{**action, "active": True} for action in _wiki_actions_metadata()
|
||||
{**action, "active": True}
|
||||
for action in _wiki_actions_metadata()
|
||||
if writes_allowed or action["name"] == "wiki_view"
|
||||
]
|
||||
if approval_required:
|
||||
_require_write_approval(entry)
|
||||
return entry
|
||||
|
||||
|
||||
@@ -513,11 +553,23 @@ def _wiki_actions_metadata() -> List[Dict[str, Any]]:
|
||||
return WikiTool().get_actions_metadata()
|
||||
|
||||
|
||||
def _require_write_approval(entry: Dict[str, Any]) -> None:
|
||||
for action in entry.get("actions") or []:
|
||||
if action.get("name") != "wiki_view":
|
||||
action["require_approval"] = True
|
||||
|
||||
|
||||
def outside_edits_allowed(source_row: Optional[Dict[str, Any]]) -> bool:
|
||||
"""Whether a wiki's owner lets API-key and widget runs edit it."""
|
||||
return bool(source_row and source_row.get("wiki_outside_edits"))
|
||||
|
||||
|
||||
def build_wiki_tool_config(
|
||||
source_id: str,
|
||||
source_owner_id: str,
|
||||
decoded_token: Optional[Dict] = None,
|
||||
user: Optional[str] = None,
|
||||
outside_caller: bool = False,
|
||||
) -> Dict[str, Any]:
|
||||
"""Build the config dict passed to the injected WikiTool."""
|
||||
return {
|
||||
@@ -525,6 +577,7 @@ def build_wiki_tool_config(
|
||||
"source_owner_id": source_owner_id,
|
||||
"decoded_token": decoded_token,
|
||||
"user": user,
|
||||
"outside_caller": bool(outside_caller),
|
||||
}
|
||||
|
||||
|
||||
@@ -534,15 +587,66 @@ def add_wiki_tool(tools_dict: Dict, config: Dict) -> None:
|
||||
Mirrors ``add_internal_search_tool``: the entry carries ``id=WIKI_TOOL_ID``
|
||||
so the executor can resolve the synthetic (DB-rowless) tool, and a ``config``
|
||||
the executor copies into the loaded tool. Mutates ``tools_dict`` in place.
|
||||
``writes_allowed=False`` (an API-key or widget run on a wiki whose owner
|
||||
hasn't allowed their edits) offers only ``wiki_view``; the tool still
|
||||
refuses writes itself from ``outside_caller``. ``approval_required`` (a
|
||||
public-link run) puts every write behind the caller's approval.
|
||||
"""
|
||||
if not config or not config.get("source_id") or not config.get("source_owner_id"):
|
||||
return
|
||||
entry = build_wiki_tool_entry()
|
||||
entry = build_wiki_tool_entry(
|
||||
writes_allowed=config.get("writes_allowed", True) is not False,
|
||||
approval_required=bool(config.get("approval_required")),
|
||||
)
|
||||
entry["id"] = WIKI_TOOL_ID
|
||||
entry["config"] = build_wiki_tool_config(
|
||||
source_id=config["source_id"],
|
||||
source_owner_id=config["source_owner_id"],
|
||||
decoded_token=config.get("decoded_token"),
|
||||
user=config.get("user"),
|
||||
outside_caller=bool(config.get("outside_caller")),
|
||||
)
|
||||
tools_dict[WIKI_TOOL_ID] = entry
|
||||
|
||||
|
||||
def apply_resume_caller_rules(
|
||||
tools_dict: Dict, *, outside_caller: bool, public_link_caller: bool
|
||||
) -> None:
|
||||
"""Hold a resumed run's saved WikiTool entry to its caller's rules.
|
||||
|
||||
A paused turn is resumed by whoever sends its tool actions, which may not
|
||||
be who started it: a widget key can resume the owner's own chat. So the
|
||||
saved entry is tightened here rather than trusted. An API-key or widget
|
||||
run gets ``outside_caller`` and, unless the wiki's owner allows their
|
||||
edits (read live), only ``wiki_view``; a public-link run approves every
|
||||
write. Never loosens an entry. Mutates ``tools_dict`` in place.
|
||||
|
||||
Args:
|
||||
tools_dict: The resumed run's tools.
|
||||
outside_caller: The saved state or the resuming request is an
|
||||
API-key or widget caller.
|
||||
public_link_caller: The saved state or the resuming request reached
|
||||
the agent through its public link.
|
||||
"""
|
||||
entry = tools_dict.get(WIKI_TOOL_ID) if isinstance(tools_dict, dict) else None
|
||||
if not isinstance(entry, dict):
|
||||
return
|
||||
if public_link_caller:
|
||||
_require_write_approval(entry)
|
||||
if not outside_caller:
|
||||
return
|
||||
config = entry.get("config")
|
||||
if not isinstance(config, dict):
|
||||
config = {}
|
||||
entry["config"] = config
|
||||
config["outside_caller"] = True
|
||||
allowed = False
|
||||
source_id = config.get("source_id")
|
||||
if source_id:
|
||||
try:
|
||||
with db_readonly() as conn:
|
||||
allowed = outside_edits_allowed(SourcesRepository(conn).get_by_id(str(source_id)))
|
||||
except Exception:
|
||||
logger.exception("Wiki outside-edits check failed for source %s", source_id)
|
||||
if not allowed:
|
||||
entry["actions"] = [a for a in entry.get("actions") or [] if a.get("name") == "wiki_view"]
|
||||
@@ -20,6 +20,8 @@ class _WorkflowNodeMixin:
|
||||
api_key: str,
|
||||
tool_ids: Optional[List[str]] = None,
|
||||
tool_principals: Optional[Dict[str, str]] = None,
|
||||
tool_owner: Optional[str] = None,
|
||||
tool_holder: Optional[dict] = None,
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__(
|
||||
@@ -35,8 +37,12 @@ class _WorkflowNodeMixin:
|
||||
# (Artifact / Code Executor / Read Document) and ``user_tools`` rows
|
||||
# alike, and an empty list means the node's LLM gets no tools.
|
||||
self.tool_executor.allowed_tool_ids = [str(t) for t in (tool_ids or [])]
|
||||
# Tools the owner can't use resolve as the editor who attached them.
|
||||
# The node's tools are the workflow owner's whoever runs it; tools the
|
||||
# owner can't use resolve as the editor who attached them.
|
||||
self.tool_executor.tool_owner = tool_owner
|
||||
self.tool_executor.tool_principals = dict(tool_principals or {})
|
||||
# The workflow row, so a dropped tool is logged with why.
|
||||
self.tool_executor.tool_holder = tool_holder
|
||||
|
||||
|
||||
class WorkflowNodeClassicAgent(_WorkflowNodeMixin, ClassicAgent):
|
||||
|
||||
@@ -405,7 +405,9 @@ class WorkflowEngine:
|
||||
"model_user_id": getattr(self.agent, "model_user_id", None),
|
||||
"api_key": node_api_key,
|
||||
"tool_ids": node_config.tools,
|
||||
"tool_owner": self._workflow_owner_id(),
|
||||
"tool_principals": self._node_tool_principals(node_config.tools),
|
||||
"tool_holder": getattr(self.agent, "workflow_row", None),
|
||||
"prompt": node_prompt,
|
||||
"chat_history": self.agent.chat_history,
|
||||
"decoded_token": self.agent.decoded_token,
|
||||
@@ -445,9 +447,9 @@ class WorkflowEngine:
|
||||
# the tool_call_attempts primary key and drop the later journal rows.
|
||||
node_executor = getattr(node_agent, "tool_executor", None)
|
||||
if node_executor is not None:
|
||||
node_executor.message_id = getattr(
|
||||
getattr(self.agent, "tool_executor", None), "message_id", None
|
||||
)
|
||||
run_executor = getattr(self.agent, "tool_executor", None)
|
||||
node_executor.message_id = getattr(run_executor, "message_id", None)
|
||||
self._inherit_caller_policy(node_executor, run_executor)
|
||||
# Run-scope the node agent's tools so artifact_generator / code_executor
|
||||
# address artifacts by this workflow run: a short ref (A1) created by one
|
||||
# node resolves for edit_artifact in a later node within the same run. Only
|
||||
@@ -1322,6 +1324,34 @@ class WorkflowEngine:
|
||||
docs_together = "\n\n".join(docs_together_parts) if docs_together_parts else None
|
||||
return docs, docs_together
|
||||
|
||||
@staticmethod
|
||||
def _inherit_caller_policy(node_executor: Any, run_executor: Any) -> None:
|
||||
"""Give a node's executor the caller rules of the run that started it.
|
||||
|
||||
A node's tools run for the same caller as the workflow agent: a
|
||||
scheduled run still can't pause, and an API-key or public-link caller
|
||||
still can't write on the owner's account unless it is allowlisted.
|
||||
|
||||
Args:
|
||||
node_executor: The node agent's ``ToolExecutor``.
|
||||
run_executor: The workflow agent's ``ToolExecutor``, if any.
|
||||
"""
|
||||
if run_executor is None:
|
||||
return
|
||||
for attr in ("headless", "external_caller", "public_link_caller"):
|
||||
setattr(node_executor, attr, bool(getattr(run_executor, attr, False)))
|
||||
for attr in ("tool_allowlist", "api_write_allowlist"):
|
||||
setattr(node_executor, attr, set(getattr(run_executor, attr, None) or ()))
|
||||
|
||||
def _workflow_owner_id(self) -> Optional[str]:
|
||||
"""The workflow's owner, whom node tools and sources run as.
|
||||
|
||||
Returns:
|
||||
The owner's user id, or None when the run has none.
|
||||
"""
|
||||
resolve_owner = getattr(self.agent, "_resolve_owner_id", None)
|
||||
return (resolve_owner() if callable(resolve_owner) else None) or self._resolve_user_id()
|
||||
|
||||
def _node_tool_principals(self, tool_ids) -> Dict[str, str]:
|
||||
"""Node tool id -> the editor to resolve it as, for sponsored tools.
|
||||
|
||||
@@ -1373,33 +1403,26 @@ class WorkflowEngine:
|
||||
if not sources:
|
||||
return []
|
||||
ids = sources if isinstance(sources, list) else [sources]
|
||||
resolve_owner = getattr(self.agent, "_resolve_owner_id", None)
|
||||
owner = (resolve_owner() if callable(resolve_owner) else None) or (
|
||||
self._resolve_user_id()
|
||||
)
|
||||
owner = self._workflow_owner_id()
|
||||
if not owner:
|
||||
logger.warning("Workflow node sources dropped: no owner to authorize.")
|
||||
return []
|
||||
|
||||
from docsgpt.api.user.resource_access import active_sponsor
|
||||
from docsgpt.api.user.team_sharing import can_access
|
||||
from docsgpt.api.user.resource_access import log_stopped, ref_access
|
||||
from docsgpt.storage.db.session import db_readonly
|
||||
|
||||
workflow_row = getattr(self.agent, "workflow_row", None)
|
||||
# The same check the workflow page's run state uses: the owner, else
|
||||
# the editor who attached it while they still qualify.
|
||||
holder = {**(getattr(self.agent, "workflow_row", None) or {}), "user_id": owner}
|
||||
allowed = []
|
||||
try:
|
||||
with db_readonly() as conn:
|
||||
for sid in ids:
|
||||
if sid and (
|
||||
can_access(conn, "source", str(sid), owner)
|
||||
or active_sponsor(conn, "workflow", workflow_row, "source", str(sid))
|
||||
):
|
||||
access = ref_access(conn, "workflow", holder, "source", str(sid)) if sid else None
|
||||
if access is not None and access.principal:
|
||||
allowed.append(sid)
|
||||
else:
|
||||
logger.warning(
|
||||
"Workflow node source %s dropped: %s has no access.",
|
||||
sid, owner,
|
||||
)
|
||||
log_stopped("workflow", holder, "source", sid, access.reason if access else None)
|
||||
except Exception:
|
||||
logger.exception("Workflow node source authorization failed; dropping all.")
|
||||
return []
|
||||
|
||||
@@ -0,0 +1,557 @@
|
||||
"""0040 connections — connector_sessions becomes the connections table.
|
||||
|
||||
``connector_sessions`` already holds one row per signed-in account (OAuth
|
||||
ingest providers) or per MCP server. This migration names what each row is
|
||||
and lets sources and tools point at the row they use:
|
||||
|
||||
* ``connector_key`` is the catalog entry (``google_drive``, ``custom_mcp``,
|
||||
``telegram``), ``auth_kind`` how the row signs in, ``display_name`` and
|
||||
``account_label`` what the Connectors page shows.
|
||||
* ``sources.connection_id`` and ``user_tools.connection_id`` link the
|
||||
resources a connection feeds. ``ON DELETE SET NULL`` keeps a source's
|
||||
indexed content when its connection is removed.
|
||||
|
||||
Connections own their credentials. Every secret of a connection (OAuth
|
||||
tokens, MCP OAuth tokens and the dynamic client registration, API keys)
|
||||
moves into ``encrypted_credentials``, a v2 envelope from
|
||||
``docsgpt.security.encryption`` bound to the owner. Plain columns
|
||||
(``status``, ``has_refresh_token``, ``scopes``, ``expires_at``) keep status
|
||||
checks from ever decrypting. The unique index gains ``account_label`` so one
|
||||
user can connect several accounts of the same service, and
|
||||
``credential_mode`` says whose account a shared tool or source uses.
|
||||
|
||||
Backfill (idempotent, only fills NULLs or unconverted rows):
|
||||
|
||||
1. ``connector_key``, ``auth_kind``, ``display_name`` from ``provider``.
|
||||
2. ``account_label`` from ``user_email`` for OAuth rows.
|
||||
3. ``sources.connection_id`` for ``connector:file`` sources, matched to the
|
||||
owner's only row for ``remote_data->>'provider'``.
|
||||
4. ``user_tools.connection_id`` for OAuth MCP tools, matched to the owner's
|
||||
row for the tool's server base URL.
|
||||
5. API-key tools (Brave, Telegram, ntfy, PostgreSQL, custom MCP with a key,
|
||||
bearer token or basic auth) get one connection per distinct credential,
|
||||
re-encrypted into v2. The tool keeps its v1 copy for one release so a
|
||||
rollback still works; the executor prefers the connection. Downgrade
|
||||
writes a v1 copy back to every tool on an API-key connection, including
|
||||
tools added after the upgrade.
|
||||
6. ``token_info`` and the secret parts of ``session_data`` (``tokens``,
|
||||
``client_info``) are encrypted into ``encrypted_credentials`` and removed
|
||||
from the plaintext columns. This needs ``ENCRYPTION_SECRET_KEY`` set to
|
||||
the value the app will run with. Downgrade decrypts them back.
|
||||
7. ``credential_mode`` is ``owner`` everywhere except OAuth MCP tools, which
|
||||
resolved each invoking member's own token before this migration and keep
|
||||
doing so (``member``); owners can switch them in the share dialog.
|
||||
|
||||
``connector_policies`` holds the admin's per-connector switches (enabled,
|
||||
forced credential mode). ``enabled`` NULL means the default: on when the
|
||||
connector has the server settings it needs, off (and hidden from members)
|
||||
until then. The instance-wide "Allow custom MCP servers" switch
|
||||
is the ``connectors.allow_custom_mcp`` key in ``app_metadata`` (absent means
|
||||
allowed).
|
||||
|
||||
Revision ID: 0040_connections
|
||||
Revises: 0039_resource_sponsors
|
||||
"""
|
||||
|
||||
from typing import Sequence, Union
|
||||
|
||||
from alembic import op
|
||||
|
||||
|
||||
revision: str = "0040_connections"
|
||||
down_revision: Union[str, None] = "0039_resource_sponsors"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
_OAUTH_PROVIDERS = {
|
||||
"google_drive": "Google Drive",
|
||||
"share_point": "SharePoint",
|
||||
"confluence": "Confluence",
|
||||
}
|
||||
|
||||
|
||||
_TOOL_CONNECTORS = {
|
||||
"brave": "Brave Search",
|
||||
"telegram": "Telegram",
|
||||
"ntfy": "ntfy",
|
||||
"postgres": "PostgreSQL",
|
||||
}
|
||||
_MCP_SECRET_AUTH = ("api_key", "bearer", "basic")
|
||||
_BATCH = 500
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
_upgrade_links()
|
||||
_upgrade_credentials()
|
||||
op.execute(
|
||||
"""
|
||||
CREATE TABLE IF NOT EXISTS connector_policies (
|
||||
connector_key TEXT PRIMARY KEY,
|
||||
enabled BOOLEAN,
|
||||
credential_mode TEXT NOT NULL DEFAULT 'choose'
|
||||
CONSTRAINT connector_policies_credential_mode_chk
|
||||
CHECK (credential_mode IN ('choose', 'owner', 'member')),
|
||||
updated_by TEXT,
|
||||
updated_at TIMESTAMPTZ NOT NULL DEFAULT now()
|
||||
);
|
||||
"""
|
||||
)
|
||||
|
||||
|
||||
def _upgrade_links() -> None:
|
||||
op.execute(
|
||||
"""
|
||||
ALTER TABLE connector_sessions
|
||||
ADD COLUMN IF NOT EXISTS connector_key TEXT,
|
||||
ADD COLUMN IF NOT EXISTS display_name TEXT,
|
||||
ADD COLUMN IF NOT EXISTS account_label TEXT,
|
||||
ADD COLUMN IF NOT EXISTS auth_kind TEXT,
|
||||
ADD COLUMN IF NOT EXISTS updated_at TIMESTAMPTZ NOT NULL DEFAULT now();
|
||||
"""
|
||||
)
|
||||
op.execute(
|
||||
"ALTER TABLE sources ADD COLUMN IF NOT EXISTS connection_id UUID "
|
||||
"REFERENCES connector_sessions(id) ON DELETE SET NULL;"
|
||||
)
|
||||
op.execute(
|
||||
"ALTER TABLE user_tools ADD COLUMN IF NOT EXISTS connection_id UUID "
|
||||
"REFERENCES connector_sessions(id) ON DELETE SET NULL;"
|
||||
)
|
||||
op.execute("CREATE INDEX IF NOT EXISTS sources_connection_idx ON sources (connection_id);")
|
||||
op.execute("CREATE INDEX IF NOT EXISTS user_tools_connection_idx ON user_tools (connection_id);")
|
||||
|
||||
# 1 + 2: name the existing rows.
|
||||
for provider, name in _OAUTH_PROVIDERS.items():
|
||||
op.execute(
|
||||
f"""
|
||||
UPDATE connector_sessions SET
|
||||
connector_key = COALESCE(connector_key, '{provider}'),
|
||||
auth_kind = COALESCE(auth_kind, 'oauth'),
|
||||
display_name = COALESCE(display_name, '{name}'),
|
||||
account_label = COALESCE(account_label, user_email)
|
||||
WHERE provider = '{provider}';
|
||||
"""
|
||||
)
|
||||
op.execute(
|
||||
"""
|
||||
UPDATE connector_sessions SET
|
||||
connector_key = COALESCE(connector_key, 'custom_mcp'),
|
||||
auth_kind = COALESCE(auth_kind, 'mcp_oauth'),
|
||||
display_name = COALESCE(
|
||||
display_name,
|
||||
regexp_replace(COALESCE(server_url, substr(provider, 5)), '^https?://', '')
|
||||
)
|
||||
WHERE provider LIKE 'mcp:%';
|
||||
"""
|
||||
)
|
||||
|
||||
# 3: connector sources point at the owner's only session for the provider.
|
||||
op.execute(
|
||||
"""
|
||||
UPDATE sources s SET connection_id = cs.id
|
||||
FROM connector_sessions cs
|
||||
WHERE s.connection_id IS NULL
|
||||
AND s.type = 'connector:file'
|
||||
AND cs.user_id = s.user_id
|
||||
AND cs.provider = s.remote_data->>'provider'
|
||||
AND (
|
||||
SELECT count(*) FROM connector_sessions c2
|
||||
WHERE c2.user_id = s.user_id AND c2.provider = cs.provider
|
||||
) = 1;
|
||||
"""
|
||||
)
|
||||
|
||||
# 4: OAuth MCP tools point at the owner's session for the server's base URL.
|
||||
op.execute(
|
||||
"""
|
||||
UPDATE user_tools t SET connection_id = cs.id
|
||||
FROM connector_sessions cs
|
||||
WHERE t.connection_id IS NULL
|
||||
AND t.name = 'mcp_tool'
|
||||
AND t.config->>'auth_type' = 'oauth'
|
||||
AND cs.user_id = t.user_id
|
||||
AND cs.provider = 'mcp:' || substring(t.config->>'server_url' from '^(https?://[^/]+)');
|
||||
"""
|
||||
)
|
||||
|
||||
|
||||
def _credential_hint(credentials: dict) -> str:
|
||||
"""``…abcd``: the last four characters of the first secret, never more."""
|
||||
for value in credentials.values():
|
||||
if isinstance(value, str) and len(value) >= 8:
|
||||
return "\u2026" + value[-4:]
|
||||
return "\u2026"
|
||||
|
||||
|
||||
def _upgrade_credentials() -> None:
|
||||
op.execute(
|
||||
"""
|
||||
ALTER TABLE connector_sessions
|
||||
ADD COLUMN IF NOT EXISTS encrypted_credentials TEXT,
|
||||
ADD COLUMN IF NOT EXISTS has_refresh_token BOOLEAN NOT NULL DEFAULT false,
|
||||
ADD COLUMN IF NOT EXISTS scopes JSONB NOT NULL DEFAULT '[]'::jsonb,
|
||||
ADD COLUMN IF NOT EXISTS last_error TEXT,
|
||||
ADD COLUMN IF NOT EXISTS last_used_at TIMESTAMPTZ;
|
||||
"""
|
||||
)
|
||||
for table in ("sources", "user_tools"):
|
||||
op.execute(
|
||||
f"ALTER TABLE {table} ADD COLUMN IF NOT EXISTS credential_mode TEXT NOT NULL DEFAULT 'owner' "
|
||||
f"CONSTRAINT {table}_credential_mode_chk CHECK (credential_mode IN ('owner', 'member'));"
|
||||
)
|
||||
op.execute("DROP INDEX IF EXISTS connector_sessions_user_endpoint_uidx;")
|
||||
op.execute(
|
||||
"CREATE UNIQUE INDEX IF NOT EXISTS connector_sessions_account_uidx ON connector_sessions "
|
||||
"(user_id, provider, COALESCE(server_url, ''), COALESCE(account_label, ''));"
|
||||
)
|
||||
|
||||
# Statuses become connected / pending / reconnect_needed / disconnected /
|
||||
# error. MCP rows never had one; name them before their tokens are sealed.
|
||||
op.execute("UPDATE connector_sessions SET status = 'connected' WHERE status = 'authorized';")
|
||||
op.execute(
|
||||
"""
|
||||
UPDATE connector_sessions SET status =
|
||||
CASE WHEN session_data ? 'tokens' THEN 'connected' ELSE 'pending' END
|
||||
WHERE provider LIKE 'mcp:%' AND status IS NULL;
|
||||
"""
|
||||
)
|
||||
|
||||
bind = op.get_bind()
|
||||
_encrypt_session_secrets(bind)
|
||||
_link_api_key_tools(bind)
|
||||
|
||||
# 7: keep OAuth MCP tools on each member's own token, as before.
|
||||
op.execute(
|
||||
"""
|
||||
UPDATE user_tools SET credential_mode = 'member'
|
||||
WHERE name = 'mcp_tool' AND config->>'auth_type' = 'oauth';
|
||||
"""
|
||||
)
|
||||
|
||||
|
||||
def _encrypt_session_secrets(bind) -> None:
|
||||
"""6: move plaintext tokens into the owner-bound v2 envelope."""
|
||||
import json
|
||||
|
||||
from sqlalchemy import text
|
||||
|
||||
from docsgpt.security.encryption import encrypt_json
|
||||
|
||||
while True:
|
||||
rows = bind.execute(
|
||||
text(
|
||||
"""
|
||||
SELECT id, user_id, token_info, token_info IS NOT NULL AS has_token_info, session_data
|
||||
FROM connector_sessions
|
||||
WHERE encrypted_credentials IS NULL
|
||||
AND (token_info IS NOT NULL OR session_data ? 'tokens' OR session_data ? 'client_info')
|
||||
LIMIT :batch
|
||||
"""
|
||||
),
|
||||
{"batch": _BATCH},
|
||||
).fetchall()
|
||||
if not rows:
|
||||
return
|
||||
for row in rows:
|
||||
token_info = _token_info_object(row.token_info)
|
||||
session_data = dict(row.session_data or {})
|
||||
secrets = {}
|
||||
if token_info is not None:
|
||||
secrets["token_info"] = token_info
|
||||
elif row.has_token_info:
|
||||
# A value the app never read as a token (a bare string, a number,
|
||||
# JSON ``null``): kept apart from ``token_info`` so downgrade can
|
||||
# restore it, and the row is not selected again.
|
||||
secrets["legacy_token_info"] = row.token_info
|
||||
for key in ("tokens", "client_info"):
|
||||
if key in session_data:
|
||||
secrets[key] = session_data.pop(key)
|
||||
tokens = secrets.get("tokens") if isinstance(secrets.get("tokens"), dict) else {}
|
||||
has_refresh = bool((token_info or {}).get("refresh_token") or tokens.get("refresh_token"))
|
||||
scopes = (token_info or {}).get("scopes") or []
|
||||
if isinstance(scopes, str):
|
||||
scopes = scopes.split()
|
||||
bind.execute(
|
||||
text(
|
||||
"""
|
||||
UPDATE connector_sessions SET
|
||||
encrypted_credentials = :blob,
|
||||
has_refresh_token = :has_refresh,
|
||||
scopes = CAST(:scopes AS jsonb),
|
||||
token_info = NULL,
|
||||
session_data = CAST(:session_data AS jsonb)
|
||||
WHERE id = :id
|
||||
"""
|
||||
),
|
||||
{
|
||||
"blob": encrypt_json(secrets, row.user_id),
|
||||
"has_refresh": has_refresh,
|
||||
"scopes": json.dumps(list(scopes)),
|
||||
"session_data": json.dumps(session_data),
|
||||
"id": row.id,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def _token_info_object(value):
|
||||
"""``token_info`` as an object, decoding one stored as a JSON string.
|
||||
|
||||
Args:
|
||||
value: The decoded ``token_info`` column value.
|
||||
|
||||
Returns:
|
||||
The object, or ``None`` when the value is not (and does not encode) one.
|
||||
"""
|
||||
import json
|
||||
|
||||
if isinstance(value, str):
|
||||
try:
|
||||
value = json.loads(value)
|
||||
except ValueError:
|
||||
return None
|
||||
return value if isinstance(value, dict) else None
|
||||
|
||||
|
||||
def _link_api_key_tools(bind) -> None:
|
||||
"""5: one connection per distinct API credential, linked from its tools."""
|
||||
from urllib.parse import urlparse
|
||||
|
||||
from sqlalchemy import text
|
||||
|
||||
from docsgpt.security.encryption import (
|
||||
CredentialDecryptionError,
|
||||
decrypt_credentials,
|
||||
decrypt_json,
|
||||
encrypt_json,
|
||||
)
|
||||
|
||||
rows = bind.execute(
|
||||
text(
|
||||
"""
|
||||
SELECT id, user_id, name, custom_name, display_name, config FROM user_tools
|
||||
WHERE connection_id IS NULL
|
||||
AND config ? 'encrypted_credentials'
|
||||
AND (
|
||||
name = ANY(:names)
|
||||
OR (name = 'mcp_tool' AND config->>'auth_type' = ANY(:mcp_auth))
|
||||
)
|
||||
"""
|
||||
),
|
||||
{"names": list(_TOOL_CONNECTORS), "mcp_auth": list(_MCP_SECRET_AUTH)},
|
||||
).fetchall()
|
||||
for row in rows:
|
||||
config = row.config or {}
|
||||
credentials = decrypt_credentials(config.get("encrypted_credentials") or "", row.user_id)
|
||||
if not credentials:
|
||||
# Written with a different key; the tool keeps its v1 copy.
|
||||
continue
|
||||
if row.name == "mcp_tool":
|
||||
parsed = urlparse(config.get("server_url") or "")
|
||||
server_url = f"{parsed.scheme}://{parsed.netloc}" if parsed.netloc else None
|
||||
connector_key = "custom_mcp"
|
||||
display_name = row.custom_name or row.display_name or parsed.netloc or "MCP server"
|
||||
else:
|
||||
server_url = None
|
||||
connector_key = row.name
|
||||
display_name = _TOOL_CONNECTORS[row.name]
|
||||
# The hint is not an identity: reuse a connection only when it holds
|
||||
# the same credentials, and give a different key its own label.
|
||||
hint = _credential_hint(credentials)
|
||||
label, suffix = hint, 1
|
||||
while True:
|
||||
existing = bind.execute(
|
||||
text(
|
||||
"""
|
||||
SELECT id, encrypted_credentials FROM connector_sessions
|
||||
WHERE user_id = :user_id AND provider = :provider
|
||||
AND COALESCE(server_url, '') = COALESCE(:server_url, '')
|
||||
AND COALESCE(account_label, '') = :label
|
||||
"""
|
||||
),
|
||||
{"user_id": row.user_id, "provider": connector_key, "server_url": server_url, "label": label},
|
||||
).fetchone()
|
||||
if existing is None:
|
||||
break
|
||||
try:
|
||||
stored = decrypt_json(existing.encrypted_credentials or "", row.user_id).get("credentials")
|
||||
except CredentialDecryptionError:
|
||||
stored = None
|
||||
if stored == credentials:
|
||||
break
|
||||
suffix += 1
|
||||
label = f"{hint} ({suffix})"
|
||||
if existing is None:
|
||||
connection_id = bind.execute(
|
||||
text(
|
||||
"""
|
||||
INSERT INTO connector_sessions (
|
||||
user_id, provider, server_url, connector_key, display_name, account_label,
|
||||
auth_kind, status, encrypted_credentials, session_data
|
||||
) VALUES (
|
||||
:user_id, :provider, :server_url, :provider, :display_name, :label,
|
||||
'api_key', 'connected', :blob, '{}'::jsonb
|
||||
) RETURNING id
|
||||
"""
|
||||
),
|
||||
{
|
||||
"user_id": row.user_id,
|
||||
"provider": connector_key,
|
||||
"server_url": server_url,
|
||||
"display_name": display_name,
|
||||
"label": label,
|
||||
"blob": encrypt_json({"credentials": credentials}, row.user_id),
|
||||
},
|
||||
).scalar()
|
||||
else:
|
||||
connection_id = existing.id
|
||||
bind.execute(
|
||||
text("UPDATE user_tools SET connection_id = :cid WHERE id = :id"),
|
||||
{"cid": connection_id, "id": row.id},
|
||||
)
|
||||
|
||||
|
||||
def _downgrade_credentials() -> None:
|
||||
"""Decrypt the envelopes back into the pre-0038 plaintext columns."""
|
||||
|
||||
from sqlalchemy import text
|
||||
|
||||
|
||||
bind = op.get_bind()
|
||||
has_envelope = bind.execute(
|
||||
text(
|
||||
"SELECT 1 FROM information_schema.columns "
|
||||
"WHERE table_name = 'connector_sessions' AND column_name = 'encrypted_credentials'"
|
||||
)
|
||||
).first()
|
||||
if has_envelope is not None:
|
||||
_decrypt_back(bind)
|
||||
_restore_account_index()
|
||||
|
||||
|
||||
def _decrypt_back(bind) -> None:
|
||||
import json
|
||||
|
||||
from sqlalchemy import text
|
||||
|
||||
from docsgpt.security.encryption import CredentialDecryptionError, decrypt_json, encrypt_credentials
|
||||
|
||||
# API-key connections go away, so their tools get a v1 copy back: tools
|
||||
# added after the upgrade never had one, and a reconnect may have changed
|
||||
# the key since the backfill.
|
||||
linked = bind.execute(
|
||||
text(
|
||||
"SELECT t.id, t.user_id, c.encrypted_credentials FROM user_tools t "
|
||||
"JOIN connector_sessions c ON c.id = t.connection_id "
|
||||
"WHERE c.auth_kind = 'api_key' AND c.encrypted_credentials IS NOT NULL"
|
||||
)
|
||||
).fetchall()
|
||||
for row in linked:
|
||||
try:
|
||||
credentials = decrypt_json(row.encrypted_credentials, row.user_id).get("credentials")
|
||||
except CredentialDecryptionError:
|
||||
continue
|
||||
if not credentials:
|
||||
continue
|
||||
bind.execute(
|
||||
text(
|
||||
"UPDATE user_tools SET config = COALESCE(config, '{}'::jsonb) "
|
||||
"|| jsonb_build_object('encrypted_credentials', CAST(:blob AS text)) WHERE id = :id"
|
||||
),
|
||||
{"blob": encrypt_credentials(credentials, row.user_id), "id": row.id},
|
||||
)
|
||||
bind.execute(text("UPDATE user_tools SET connection_id = NULL WHERE connection_id IN "
|
||||
"(SELECT id FROM connector_sessions WHERE auth_kind = 'api_key')"))
|
||||
bind.execute(text("DELETE FROM connector_sessions WHERE auth_kind = 'api_key'"))
|
||||
rows = bind.execute(
|
||||
text(
|
||||
"SELECT id, user_id, session_data, encrypted_credentials FROM connector_sessions "
|
||||
"WHERE encrypted_credentials IS NOT NULL"
|
||||
)
|
||||
).fetchall()
|
||||
for row in rows:
|
||||
try:
|
||||
secrets = decrypt_json(row.encrypted_credentials, row.user_id)
|
||||
except CredentialDecryptionError:
|
||||
continue
|
||||
session_data = dict(row.session_data or {})
|
||||
for key in ("tokens", "client_info"):
|
||||
if key in secrets:
|
||||
session_data[key] = secrets[key]
|
||||
bind.execute(
|
||||
text(
|
||||
"UPDATE connector_sessions SET token_info = CAST(:token_info AS jsonb), "
|
||||
"session_data = CAST(:session_data AS jsonb) WHERE id = :id"
|
||||
),
|
||||
{
|
||||
"token_info": _restored_token_info(secrets),
|
||||
"session_data": json.dumps(session_data),
|
||||
"id": row.id,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def _restored_token_info(secrets: dict):
|
||||
"""The ``token_info`` JSON downgrade writes back, ``None`` for SQL NULL.
|
||||
|
||||
Args:
|
||||
secrets: The decrypted envelope.
|
||||
|
||||
Returns:
|
||||
The JSON text of ``token_info`` or of a preserved non-object value.
|
||||
"""
|
||||
import json
|
||||
|
||||
for key in ("token_info", "legacy_token_info"):
|
||||
if key in secrets:
|
||||
return json.dumps(secrets[key])
|
||||
return None
|
||||
|
||||
|
||||
def _restore_account_index() -> None:
|
||||
# Several accounts per provider cannot survive the old unique index: keep
|
||||
# the most recently updated one.
|
||||
op.execute(
|
||||
"""
|
||||
DELETE FROM connector_sessions c USING connector_sessions newer
|
||||
WHERE c.user_id = newer.user_id AND c.provider = newer.provider
|
||||
AND COALESCE(c.server_url, '') = COALESCE(newer.server_url, '')
|
||||
AND (c.updated_at, c.id) < (newer.updated_at, newer.id);
|
||||
"""
|
||||
)
|
||||
op.execute("DROP INDEX IF EXISTS connector_sessions_account_uidx;")
|
||||
op.execute(
|
||||
"CREATE UNIQUE INDEX IF NOT EXISTS connector_sessions_user_endpoint_uidx "
|
||||
"ON connector_sessions (user_id, COALESCE(server_url, ''), provider);"
|
||||
)
|
||||
for table in ("sources", "user_tools"):
|
||||
op.execute(f"ALTER TABLE {table} DROP COLUMN IF EXISTS credential_mode;")
|
||||
op.execute(
|
||||
"""
|
||||
ALTER TABLE connector_sessions
|
||||
DROP COLUMN IF EXISTS last_used_at,
|
||||
DROP COLUMN IF EXISTS last_error,
|
||||
DROP COLUMN IF EXISTS scopes,
|
||||
DROP COLUMN IF EXISTS has_refresh_token,
|
||||
DROP COLUMN IF EXISTS encrypted_credentials;
|
||||
"""
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.execute("DROP TABLE IF EXISTS connector_policies;")
|
||||
_downgrade_credentials()
|
||||
op.execute("DROP INDEX IF EXISTS user_tools_connection_idx;")
|
||||
op.execute("DROP INDEX IF EXISTS sources_connection_idx;")
|
||||
op.execute("ALTER TABLE user_tools DROP COLUMN IF EXISTS connection_id;")
|
||||
op.execute("ALTER TABLE sources DROP COLUMN IF EXISTS connection_id;")
|
||||
op.execute(
|
||||
"""
|
||||
ALTER TABLE connector_sessions
|
||||
DROP COLUMN IF EXISTS updated_at,
|
||||
DROP COLUMN IF EXISTS auth_kind,
|
||||
DROP COLUMN IF EXISTS account_label,
|
||||
DROP COLUMN IF EXISTS display_name,
|
||||
DROP COLUMN IF EXISTS connector_key;
|
||||
"""
|
||||
)
|
||||
@@ -0,0 +1,32 @@
|
||||
"""0041 connection account name — what the user calls an account.
|
||||
|
||||
``account_label`` identifies an account: the email an OAuth sign-in returns,
|
||||
or a hint of a pasted key. Signing in again finds the connection by it, so it
|
||||
cannot be renamed. ``account_name`` is the name the user gives the account
|
||||
("Alerts bot", "Work"), shown instead of the label and used to tell two
|
||||
accounts of one service apart, for people and for the model. NULL means the
|
||||
user never named it.
|
||||
|
||||
Idempotent both ways.
|
||||
|
||||
Revision ID: 0041_connection_account_name
|
||||
Revises: 0040_connections
|
||||
"""
|
||||
|
||||
from typing import Sequence, Union
|
||||
|
||||
from alembic import op
|
||||
|
||||
|
||||
revision: str = "0041_connection_account_name"
|
||||
down_revision: Union[str, None] = "0040_connections"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.execute("ALTER TABLE connector_sessions ADD COLUMN IF NOT EXISTS account_name TEXT;")
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.execute("ALTER TABLE connector_sessions DROP COLUMN IF EXISTS account_name;")
|
||||
@@ -0,0 +1,40 @@
|
||||
"""0042 schedule created_via api — schedules set by an API-key caller.
|
||||
|
||||
A schedule the agent sets from a widget or API chat is stored under the
|
||||
agent's owner, like the run itself. ``created_via = 'api'`` records that
|
||||
it came from outside the app, so its runs keep that caller's limits: no
|
||||
writes on the owner's accounts or credentials unless the owner allowed
|
||||
them in the agent's API write allowlist.
|
||||
|
||||
Idempotent both ways. Downgrade folds ``api`` back into ``chat``.
|
||||
|
||||
Revision ID: 0042_schedule_created_via_api
|
||||
Revises: 0041_connection_account_name
|
||||
"""
|
||||
|
||||
from typing import Sequence, Union
|
||||
|
||||
from alembic import op
|
||||
|
||||
|
||||
revision: str = "0042_schedule_created_via_api"
|
||||
down_revision: Union[str, None] = "0041_connection_account_name"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.execute("ALTER TABLE schedules DROP CONSTRAINT IF EXISTS schedules_created_via_chk;")
|
||||
op.execute(
|
||||
"ALTER TABLE schedules ADD CONSTRAINT schedules_created_via_chk "
|
||||
"CHECK (created_via IN ('chat', 'ui', 'api'));"
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.execute("ALTER TABLE schedules DROP CONSTRAINT IF EXISTS schedules_created_via_chk;")
|
||||
op.execute("UPDATE schedules SET created_via = 'chat' WHERE created_via = 'api';")
|
||||
op.execute(
|
||||
"ALTER TABLE schedules ADD CONSTRAINT schedules_created_via_chk "
|
||||
"CHECK (created_via IN ('chat', 'ui'));"
|
||||
)
|
||||
@@ -0,0 +1,32 @@
|
||||
"""0043 wiki outside edits — the wiki owner's say on API and widget edits.
|
||||
|
||||
An agent run from its API key or widget acts as the agent's owner, so it
|
||||
could rewrite any wiki the owner can edit. ``wiki_outside_edits`` records
|
||||
whether the wiki's owner allows that; it is off by default, so such runs can
|
||||
still read the wiki but not change it.
|
||||
|
||||
Idempotent both ways.
|
||||
|
||||
Revision ID: 0043_wiki_outside_edits
|
||||
Revises: 0042_schedule_created_via_api
|
||||
"""
|
||||
|
||||
from typing import Sequence, Union
|
||||
|
||||
from alembic import op
|
||||
|
||||
|
||||
revision: str = "0043_wiki_outside_edits"
|
||||
down_revision: Union[str, None] = "0042_schedule_created_via_api"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.execute(
|
||||
"ALTER TABLE sources ADD COLUMN IF NOT EXISTS wiki_outside_edits BOOLEAN NOT NULL DEFAULT false;"
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.execute("ALTER TABLE sources DROP COLUMN IF EXISTS wiki_outside_edits;")
|
||||
@@ -1,5 +1,6 @@
|
||||
from .routes import admin_ns
|
||||
from . import quotas # noqa: F401 (registers the quota resources on admin_ns)
|
||||
from . import activity # noqa: F401 (registers the activity resources on admin_ns)
|
||||
from . import connectors # noqa: F401 (registers the connector policy resources on admin_ns)
|
||||
|
||||
__all__ = ["admin_ns"]
|
||||
@@ -0,0 +1,159 @@
|
||||
"""Admin > Connectors: turn connectors on or off and set sharing policy.
|
||||
|
||||
Every resource here is behind ``@admin_required`` (the frontend guard is
|
||||
cosmetic). The page also shows what each OAuth connector still needs from
|
||||
the server (settings and the redirect URI to register), never their values.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
|
||||
from flask import jsonify, make_response, request
|
||||
from flask_restx import Resource
|
||||
from sqlalchemy import text
|
||||
|
||||
from docsgpt.api.admin.routes import _actor, admin_ns
|
||||
from docsgpt.api.user.authz import admin_required
|
||||
from docsgpt.connectors import catalog, service
|
||||
from docsgpt.core.settings import settings
|
||||
from docsgpt.security.encryption import is_default_encryption_key
|
||||
from docsgpt.storage.db.repositories.app_metadata import AppMetadataRepository
|
||||
from docsgpt.storage.db.repositories.connector_policies import (
|
||||
ALLOW_CUSTOM_MCP_KEY,
|
||||
CREDENTIAL_POLICIES,
|
||||
ConnectorPoliciesRepository,
|
||||
allow_writes_key,
|
||||
)
|
||||
from docsgpt.storage.db.session import db_readonly, db_session
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _connection_counts(conn) -> dict[str, int]:
|
||||
rows = conn.execute(
|
||||
text(
|
||||
"SELECT connector_key, provider, server_url, count(*) AS n FROM connector_sessions "
|
||||
"WHERE COALESCE(status, '') <> 'pending' GROUP BY connector_key, provider, server_url"
|
||||
)
|
||||
).fetchall()
|
||||
counts: dict[str, int] = {}
|
||||
for row in rows:
|
||||
key = catalog.connector_key_for_row(dict(row._mapping))
|
||||
if key:
|
||||
counts[key] = counts.get(key, 0) + int(row.n)
|
||||
return counts
|
||||
|
||||
|
||||
def _mcp_redirect_uri() -> str:
|
||||
from docsgpt.agents.tools.mcp_tool import MCPTool
|
||||
|
||||
return MCPTool._resolve_redirect_uri(MCPTool.__new__(MCPTool), None)
|
||||
|
||||
|
||||
@admin_ns.route("/admin/connectors")
|
||||
class AdminConnectorsResource(Resource):
|
||||
@admin_required
|
||||
def get(self):
|
||||
"""Every connector with its policy, setup state and connection count."""
|
||||
with db_readonly() as conn:
|
||||
policies = ConnectorPoliciesRepository(conn).all()
|
||||
loaded = service.load_policies(conn)
|
||||
allow_custom = service.custom_mcp_allowed(conn)
|
||||
counts = _connection_counts(conn)
|
||||
connectors = []
|
||||
for definition in catalog.all_definitions():
|
||||
policy = policies.get(definition.key) or {}
|
||||
connectors.append(
|
||||
{
|
||||
"key": definition.key,
|
||||
"name": definition.name,
|
||||
"icon": definition.icon,
|
||||
"category": definition.category,
|
||||
"publisher": definition.publisher,
|
||||
"auth_kind": definition.auth_kind,
|
||||
"capabilities": list(definition.capabilities),
|
||||
"enabled": service.connector_is_enabled(policies, definition.key),
|
||||
"credential_mode": policy.get("credential_mode", "choose"),
|
||||
"configured": definition.configured,
|
||||
"required_settings": [
|
||||
{"name": name, "set": bool(getattr(settings, name, None))}
|
||||
for name in definition.required_settings
|
||||
],
|
||||
# Optional: they add a second sign-in (GitHub's App) to a
|
||||
# connector that already works without them.
|
||||
"oauth_settings": [
|
||||
{"name": name, "set": bool(getattr(settings, name, None))}
|
||||
for name in definition.oauth_settings
|
||||
],
|
||||
"oauth_configured": definition.oauth_configured,
|
||||
"connection_count": counts.get(definition.key, 0),
|
||||
"docs_url": definition.docs_url,
|
||||
"mcp_url": definition.mcp_url,
|
||||
# Whether members may let agents make changes (GitHub);
|
||||
# None where the connector offers no such choice.
|
||||
"allow_writes": (
|
||||
service.writes_allowed(loaded, definition.key) if definition.mcp_write_url else None
|
||||
),
|
||||
}
|
||||
)
|
||||
return make_response(
|
||||
jsonify(
|
||||
{
|
||||
"success": True,
|
||||
"connectors": connectors,
|
||||
"allow_custom_mcp": allow_custom,
|
||||
"default_encryption_key": is_default_encryption_key(),
|
||||
"oauth_redirect_uri": settings.CONNECTOR_REDIRECT_BASE_URI,
|
||||
"mcp_redirect_uri": _mcp_redirect_uri(),
|
||||
}
|
||||
),
|
||||
200,
|
||||
)
|
||||
|
||||
@admin_required
|
||||
def put(self):
|
||||
"""Update policies: {policies: {key: {enabled?, credential_mode?, allow_writes?}}, allow_custom_mcp?}."""
|
||||
body = request.get_json(silent=True) or {}
|
||||
updates = body.get("policies") or {}
|
||||
if not isinstance(updates, dict):
|
||||
return make_response(jsonify({"success": False, "message": "policies must be an object"}), 400)
|
||||
for key, change in updates.items():
|
||||
if catalog.get_definition(key) is None or not isinstance(change, dict):
|
||||
return make_response(jsonify({"success": False, "message": f"Unknown connector: {key}"}), 400)
|
||||
mode = change.get("credential_mode")
|
||||
if mode is not None and mode not in CREDENTIAL_POLICIES:
|
||||
return make_response(jsonify({"success": False, "message": "Unknown credential mode"}), 400)
|
||||
if "enabled" in change and not isinstance(change["enabled"], bool):
|
||||
return make_response(jsonify({"success": False, "message": "enabled must be true or false"}), 400)
|
||||
if "allow_writes" in change:
|
||||
if not catalog.get_definition(key).mcp_write_url:
|
||||
return make_response(
|
||||
jsonify({"success": False, "message": f"{key} has no write access to allow"}), 400,
|
||||
)
|
||||
if not isinstance(change["allow_writes"], bool):
|
||||
return make_response(
|
||||
jsonify({"success": False, "message": "allow_writes must be true or false"}), 400,
|
||||
)
|
||||
allow_custom = body.get("allow_custom_mcp")
|
||||
if allow_custom is not None and not isinstance(allow_custom, bool):
|
||||
return make_response(jsonify({"success": False, "message": "allow_custom_mcp must be a boolean"}), 400)
|
||||
actor = _actor()
|
||||
with db_session() as conn:
|
||||
repo = ConnectorPoliciesRepository(conn)
|
||||
for key, change in updates.items():
|
||||
if "allow_writes" in change:
|
||||
value = "true" if change["allow_writes"] else "false"
|
||||
AppMetadataRepository(conn).set(allow_writes_key(key), value)
|
||||
if set(change) == {"allow_writes"}:
|
||||
continue
|
||||
repo.upsert(
|
||||
key,
|
||||
enabled=change.get("enabled"),
|
||||
credential_mode=change.get("credential_mode"),
|
||||
updated_by=actor,
|
||||
)
|
||||
if allow_custom is not None:
|
||||
AppMetadataRepository(conn).set(ALLOW_CUSTOM_MCP_KEY, "true" if allow_custom else "false")
|
||||
logger.info("connector_policies_updated", extra={"admin": actor, "connectors": sorted(updates)})
|
||||
return self.get()
|
||||
@@ -1057,6 +1057,17 @@ class BaseAnswerResource:
|
||||
"llm_name": getattr(agent, "llm_name", settings.LLM_PROVIDER),
|
||||
"api_key": getattr(agent, "api_key", None),
|
||||
"user_api_key": user_api_key,
|
||||
# An API-key caller stays one after a
|
||||
# resume (owner-account write rules).
|
||||
"external_api_caller": getattr(
|
||||
agent.tool_executor, "external_caller", False,
|
||||
),
|
||||
"public_link_caller": getattr(
|
||||
agent.tool_executor, "public_link_caller", False,
|
||||
),
|
||||
"api_write_allowlist": sorted(
|
||||
getattr(agent.tool_executor, "api_write_allowlist", set()),
|
||||
),
|
||||
"agent_id": agent_id,
|
||||
"agent_type": agent.__class__.__name__,
|
||||
"prompt": getattr(agent, "prompt", ""),
|
||||
|
||||
@@ -4,13 +4,12 @@ import json
|
||||
import logging
|
||||
import uuid
|
||||
from pathlib import Path
|
||||
from typing import Any, Callable, Dict, Optional, Set, TypeVar
|
||||
from typing import Any, Callable, Dict, List, Optional, Set, Tuple, TypeVar
|
||||
|
||||
from flask import after_this_request
|
||||
|
||||
from docsgpt import tracing
|
||||
from docsgpt.agents.agent_creator import AgentCreator
|
||||
from docsgpt.agents.default_tools import synthesized_default_tools
|
||||
from docsgpt.api.answer.services.compression import CompressionOrchestrator
|
||||
from docsgpt.api.answer.services.compression.token_counter import TokenCounter
|
||||
from docsgpt.api.answer.services.compression.types import is_compression_summary_row
|
||||
@@ -28,7 +27,9 @@ from docsgpt.core.model_utils import (
|
||||
get_provider_from_model_id,
|
||||
validate_model_id,
|
||||
)
|
||||
from docsgpt.agents.tools.wiki import apply_resume_caller_rules, outside_edits_allowed
|
||||
from docsgpt.core.settings import settings
|
||||
from docsgpt.guardrails.config import AgentConfig
|
||||
from sqlalchemy import text as sql_text
|
||||
|
||||
from docsgpt.storage.db.base_repository import looks_like_uuid, row_to_dict
|
||||
@@ -37,8 +38,6 @@ from docsgpt.storage.db.repositories.attachments import AttachmentsRepository
|
||||
from docsgpt.storage.db.repositories.prompts import PromptsRepository
|
||||
from docsgpt.storage.db.repositories.sources import SourcesRepository
|
||||
from docsgpt.storage.db.repositories.team_scope import TeamScopeRepository
|
||||
from docsgpt.storage.db.repositories.user_tools import UserToolsRepository
|
||||
from docsgpt.storage.db.repositories.users import UsersRepository
|
||||
from docsgpt.api.user.team_sharing import can_access
|
||||
from docsgpt.storage.db.session import db_readonly, db_session
|
||||
from docsgpt.storage.db.source_config import SourceConfig
|
||||
@@ -52,6 +51,25 @@ from docsgpt.utils import (
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def is_external_api_caller(data: Dict[str, Any], decoded_token: Optional[Dict], owner: Optional[str]) -> bool:
|
||||
"""Whether a request calls an agent with its API key on someone else's behalf.
|
||||
|
||||
Widget and API requests carry the agent's key and run as its owner. The
|
||||
owner previewing their own agent in the app sends the key too, but is
|
||||
signed in as that owner. In local mode without auth everyone is the same
|
||||
user, so there is no one else to tell apart.
|
||||
|
||||
Args:
|
||||
data: The request body.
|
||||
decoded_token: The caller's token before the key's owner replaces it.
|
||||
owner: The agent owner's user id.
|
||||
"""
|
||||
if not data.get("api_key"):
|
||||
return False
|
||||
caller = (decoded_token or {}).get("sub")
|
||||
return not caller or caller != owner
|
||||
|
||||
|
||||
def _clamp_chunks(value: int) -> int:
|
||||
"""Bound top-k to the range ``RetrievalConfig`` enforces, keeping 0.
|
||||
|
||||
@@ -111,6 +129,10 @@ def get_prompt(prompt_id: str, prompts_collection=None) -> str:
|
||||
|
||||
_PROMPT_PRESETS_WITHOUT_ROW = ("reduce",)
|
||||
|
||||
# Tools whose approval is decided per call from live state, not the stored
|
||||
# ``require_approval`` flags (see ``ToolExecutor.check_pause``).
|
||||
_LIVE_APPROVAL_TOOLS = frozenset({"remote_device", "code_executor"})
|
||||
|
||||
|
||||
def authorized_prompt_id(prompt_id: Any, principal: Optional[str], agent: Optional[dict] = None) -> Any:
|
||||
"""``prompt_id`` if ``principal`` (or the agent's sponsor) may use it, else ``"default"``.
|
||||
@@ -134,20 +156,27 @@ def authorized_prompt_id(prompt_id: Any, principal: Optional[str], agent: Option
|
||||
pid = str(prompt_id)
|
||||
if is_composed_preset(pid) or pid in _PROMPT_PRESETS_WITHOUT_ROW:
|
||||
return prompt_id
|
||||
from docsgpt.api.user.resource_access import active_sponsor, resolve
|
||||
from docsgpt.api.user.resource_access import (
|
||||
REASON_OWNER_LOST_ACCESS,
|
||||
log_stopped,
|
||||
ref_access,
|
||||
)
|
||||
|
||||
# The same check the agent page's run state uses: the principal, else a
|
||||
# live sponsor on the agent.
|
||||
holder = {**agent, "user_id": principal} if agent and agent.get("id") else {"user_id": principal}
|
||||
try:
|
||||
with db_readonly() as conn:
|
||||
ra = resolve(conn, "prompt", pid, principal) if principal else None
|
||||
usable = ra is not None and ra.can("use")
|
||||
if not usable and agent and agent.get("id"):
|
||||
usable = active_sponsor(conn, "agent", agent, "prompt", pid) is not None
|
||||
access = ref_access(conn, "agent", holder, "prompt", pid) if principal else None
|
||||
except Exception:
|
||||
logger.exception("Prompt access check failed for %s", pid)
|
||||
usable = False
|
||||
if usable:
|
||||
access = None
|
||||
if access is not None and access.principal:
|
||||
return prompt_id
|
||||
logger.info("prompt %s not usable by %s; using the default prompt", pid, principal)
|
||||
log_stopped(
|
||||
"agent" if agent else "chat", holder, "prompt", pid,
|
||||
access.reason if access is not None else REASON_OWNER_LOST_ACCESS,
|
||||
)
|
||||
return "default"
|
||||
|
||||
|
||||
@@ -158,14 +187,48 @@ def _agent_source_doc(conn: Any, sources_repo: Any, agent: dict, source_id: Any)
|
||||
editor who attached it while they still qualify. Read unscoped once
|
||||
authorized: an owner-scoped read misses a team-shared source.
|
||||
"""
|
||||
from docsgpt.api.user.resource_access import ref_principal
|
||||
from docsgpt.api.user.resource_access import log_stopped, ref_access
|
||||
|
||||
if not ref_principal(conn, "agent", agent, "source", str(source_id)):
|
||||
logger.info("agent %s source %s not usable; skipped", agent.get("id"), source_id)
|
||||
access = ref_access(conn, "agent", agent, "source", str(source_id))
|
||||
if not access.principal:
|
||||
log_stopped("agent", agent, "source", source_id, access.reason)
|
||||
return None
|
||||
return sources_repo.get_by_id(str(source_id))
|
||||
|
||||
|
||||
def authorized_agent_sources(conn: Any, agent: dict) -> Tuple[Optional[dict], List[dict]]:
|
||||
"""The source rows an agent run retrieves from: primary first, then extras.
|
||||
|
||||
Each is authorized like :func:`_agent_source_doc` (the owner, else the
|
||||
editor who attached it), and a source listed twice appears once.
|
||||
|
||||
Args:
|
||||
conn: An open database connection.
|
||||
agent: The ``agents`` row.
|
||||
|
||||
Returns:
|
||||
The primary source row (None when unset or not usable) and every
|
||||
usable row in run order.
|
||||
"""
|
||||
sources_repo = SourcesRepository(conn)
|
||||
primary: Optional[dict] = None
|
||||
rows: List[dict] = []
|
||||
seen: set = set()
|
||||
refs = [(True, agent.get("source_id"))]
|
||||
refs.extend((False, sid) for sid in agent.get("extra_source_ids") or [])
|
||||
for is_primary, sid_raw in refs:
|
||||
if not sid_raw:
|
||||
continue
|
||||
source_doc = _agent_source_doc(conn, sources_repo, agent, sid_raw)
|
||||
if not source_doc or str(source_doc["id"]) in seen:
|
||||
continue
|
||||
if is_primary:
|
||||
primary = source_doc
|
||||
seen.add(str(source_doc["id"]))
|
||||
rows.append(source_doc)
|
||||
return primary, rows
|
||||
|
||||
|
||||
def _wiki_write_owner(conn: Any, source_id: str, caller: str) -> Optional[str]:
|
||||
"""The owner id to write a wiki source as, when ``caller`` may edit it."""
|
||||
from docsgpt.api.user.resource_access import resolve
|
||||
@@ -231,12 +294,26 @@ class StreamProcessor:
|
||||
request_data: Dict[str, Any],
|
||||
decoded_token: Optional[Dict[str, Any]],
|
||||
trace_source: str = "stream",
|
||||
*,
|
||||
external_caller: bool = False,
|
||||
):
|
||||
"""Bind a request to its processor.
|
||||
|
||||
Args:
|
||||
request_data: The request body.
|
||||
decoded_token: The caller's token; ``/v1`` passes the agent owner's.
|
||||
trace_source: The entry point the trace is stored under.
|
||||
external_caller: Set by the server for requests authenticated by
|
||||
an agent's API key whose token is the owner's (``/v1``), so
|
||||
the run keeps the key holder's write limits. Never read from
|
||||
the request body.
|
||||
"""
|
||||
# Legacy attribute retained as None for any external callers that
|
||||
# introspect the processor; all DB access uses per-op connections.
|
||||
self.prompts_collection = None
|
||||
self.data = request_data
|
||||
self.decoded_token = decoded_token
|
||||
self.external_caller = bool(external_caller)
|
||||
self.initial_user_id = (
|
||||
self.decoded_token.get("sub") if self.decoded_token is not None else None
|
||||
)
|
||||
@@ -250,6 +327,9 @@ class StreamProcessor:
|
||||
self.retriever_config = {}
|
||||
self.is_shared_usage = False
|
||||
self.shared_token = None
|
||||
# Set by _get_agent_key: the caller reaches the agent only through its
|
||||
# public link (not its owner, no team grant).
|
||||
self.public_link_usage = False
|
||||
self.agent_id = self.data.get("agent_id")
|
||||
# Set by _get_agent_key once access checks pass; read for keyless runs.
|
||||
self._authorized_agent_row: Optional[Dict[str, Any]] = None
|
||||
@@ -325,7 +405,13 @@ class StreamProcessor:
|
||||
share it with the rest of the turn. It is always generated here, never
|
||||
taken from the request body: request quotas count distinct request
|
||||
ids, so a client-chosen id would let every call count as one.
|
||||
|
||||
A request with neither a token nor an agent API key has nobody to run
|
||||
for: nothing is set up (no pre-fetch runs the agent's tools), None is
|
||||
returned and the route answers 401.
|
||||
"""
|
||||
if not self.decoded_token and not self.data.get("api_key"):
|
||||
return None
|
||||
if not getattr(self, "request_id", None):
|
||||
self.request_id = str(uuid.uuid4())
|
||||
self.initialize()
|
||||
@@ -682,9 +768,10 @@ class StreamProcessor:
|
||||
# Team-shared agents are runnable by any member with a grant
|
||||
# (viewer is enough to run). Resolved live against team_members
|
||||
# on the SAME connection so a revoked grant/membership denies on
|
||||
# the next call; resolution failure fails closed.
|
||||
# the next call; resolution failure fails closed. Checked on a
|
||||
# public agent too: a teammate there is not a link user.
|
||||
is_team_shared = False
|
||||
if not (is_owner or is_shared_with_user) and user_id:
|
||||
if not is_owner and user_id:
|
||||
try:
|
||||
is_team_shared = TeamScopeRepository(conn).can_read(
|
||||
user_id, "agent", str(agent["id"])
|
||||
@@ -697,6 +784,7 @@ class StreamProcessor:
|
||||
|
||||
if not (is_owner or is_shared_with_user or is_team_shared):
|
||||
raise Exception("Unauthorized access to the agent")
|
||||
self.public_link_usage = not (is_owner or is_team_shared)
|
||||
# Authorized. Keep the row so _configure_agent can read fields that
|
||||
# do not depend on an API key — a draft agent has key = NULL, and
|
||||
# the builder preview runs exactly that path.
|
||||
@@ -729,7 +817,6 @@ class StreamProcessor:
|
||||
agent = AgentsRepository(conn).find_by_key(api_key)
|
||||
if not agent:
|
||||
raise Exception("Invalid API Key, please generate a new key", 401)
|
||||
sources_repo = SourcesRepository(conn)
|
||||
# The repo dict uses "user_id" — the streaming path expects
|
||||
# a "user" key (legacy Mongo shape) for identity propagation.
|
||||
data: Dict[str, Any] = dict(agent)
|
||||
@@ -739,68 +826,29 @@ class StreamProcessor:
|
||||
# ``_configure_source`` ignores an empty ``data["sources"]``,
|
||||
# so the primary must appear in the union too — not only in
|
||||
# the legacy ``data["source"]`` slot.
|
||||
sources_list: list = []
|
||||
seen: set = set()
|
||||
primary_id = agent.get("source_id")
|
||||
# ``sources`` row may have NULL ``retriever``/``chunks`` —
|
||||
# fall back to the agent's value (``dict.get`` returns None
|
||||
# even when the key exists with value None).
|
||||
if primary_id:
|
||||
source_doc = _agent_source_doc(conn, sources_repo, agent, primary_id)
|
||||
if source_doc:
|
||||
sid = str(source_doc["id"])
|
||||
data["source"] = sid
|
||||
src_retriever = source_doc.get("retriever")
|
||||
if src_retriever:
|
||||
data["retriever"] = src_retriever
|
||||
src_chunks = source_doc.get("chunks")
|
||||
if src_chunks is not None:
|
||||
data["chunks"] = src_chunks
|
||||
sources_list.append(
|
||||
{
|
||||
"id": sid,
|
||||
"retriever": src_retriever or "classic",
|
||||
"chunks": (
|
||||
src_chunks if src_chunks is not None
|
||||
else data.get("chunks", "6")
|
||||
),
|
||||
# Per-source behaviour contract (lenient read).
|
||||
"retrieval": SourceConfig.parse(
|
||||
source_doc.get("config")
|
||||
).retrieval,
|
||||
}
|
||||
)
|
||||
seen.add(sid)
|
||||
else:
|
||||
data["source"] = None
|
||||
else:
|
||||
data["source"] = None
|
||||
|
||||
for sid_raw in agent.get("extra_source_ids") or []:
|
||||
if not sid_raw:
|
||||
continue
|
||||
source_doc = _agent_source_doc(conn, sources_repo, agent, sid_raw)
|
||||
if not source_doc:
|
||||
continue
|
||||
sid = str(source_doc["id"])
|
||||
if sid in seen:
|
||||
continue
|
||||
src_retriever = source_doc.get("retriever")
|
||||
src_chunks = source_doc.get("chunks")
|
||||
sources_list.append(
|
||||
{
|
||||
"id": sid,
|
||||
"retriever": src_retriever or "classic",
|
||||
"chunks": (
|
||||
src_chunks if src_chunks is not None
|
||||
else data.get("chunks", "6")
|
||||
),
|
||||
"retrieval": SourceConfig.parse(
|
||||
source_doc.get("config")
|
||||
).retrieval,
|
||||
}
|
||||
)
|
||||
seen.add(sid)
|
||||
primary, source_docs = authorized_agent_sources(conn, agent)
|
||||
# ``sources`` row may have NULL ``retriever``/``chunks`` — fall back to
|
||||
# the agent's value (``dict.get`` returns None even when the key
|
||||
# exists with value None). The primary's own values win for the agent.
|
||||
data["source"] = str(primary["id"]) if primary else None
|
||||
if primary:
|
||||
if primary.get("retriever"):
|
||||
data["retriever"] = primary["retriever"]
|
||||
if primary.get("chunks") is not None:
|
||||
data["chunks"] = primary["chunks"]
|
||||
sources_list: list = [
|
||||
{
|
||||
"id": str(source_doc["id"]),
|
||||
"retriever": source_doc.get("retriever") or "classic",
|
||||
"chunks": (
|
||||
source_doc["chunks"] if source_doc.get("chunks") is not None
|
||||
else data.get("chunks", "6")
|
||||
),
|
||||
# Per-source behaviour contract (lenient read).
|
||||
"retrieval": SourceConfig.parse(source_doc.get("config")).retrieval,
|
||||
}
|
||||
for source_doc in source_docs
|
||||
]
|
||||
data["sources"] = sources_list
|
||||
data["default_model_id"] = data.get("default_model_id", "")
|
||||
return data
|
||||
@@ -987,6 +1035,9 @@ class StreamProcessor:
|
||||
agent_id, self.initial_user_id
|
||||
)
|
||||
self.agent_id = str(agent_id) if agent_id else None
|
||||
self.agent_config["public_link_caller"] = bool(
|
||||
self.agent_id and getattr(self, "public_link_usage", False)
|
||||
)
|
||||
|
||||
# Determine the effective API key (explicit > agent-derived)
|
||||
effective_key = self.data.get("api_key") or self.agent_key
|
||||
@@ -1024,6 +1075,13 @@ class StreamProcessor:
|
||||
)
|
||||
|
||||
# Set identity context
|
||||
owner = self._agent_data.get("user")
|
||||
self.agent_config["external_api_caller"] = getattr(self, "external_caller", False) or (
|
||||
is_external_api_caller(self.data, self.decoded_token, owner)
|
||||
)
|
||||
self.agent_config["api_write_allowlist"] = AgentConfig.parse(
|
||||
self._agent_data.get("config")
|
||||
).api_write_allowlist
|
||||
if self.data.get("api_key"):
|
||||
# External API key: use the key owner's identity
|
||||
self.initial_user_id = self._agent_data.get("user")
|
||||
@@ -1210,10 +1268,25 @@ class StreamProcessor:
|
||||
writable wiki source; the first match wins and the scan stops there so
|
||||
this runs at most one owner+source lookup per chat on the hot path.
|
||||
Returns None when no writable wiki source is present.
|
||||
|
||||
An API-key or widget run (``outside_caller``) acts as the agent's
|
||||
owner, so it gets the edit actions only when the wiki's owner turned
|
||||
on ``wiki_outside_edits``; otherwise ``writes_allowed`` is False and
|
||||
the tool offers only ``wiki_view``. A public-link visitor runs as
|
||||
themselves, so they reach only wikis they may edit anyway; the switch
|
||||
doesn't apply to them, but each of their edits waits for their
|
||||
approval (``approval_required``), so the agent's prompt or sources
|
||||
can't steer the model into changing their wiki unasked.
|
||||
"""
|
||||
caller = self.decoded_token.get("sub") if self.decoded_token else None
|
||||
if not caller:
|
||||
return None
|
||||
# Processors built without __init__ (tests, resume helpers) lack these.
|
||||
run_config = getattr(self, "agent_config", None) or {}
|
||||
outside_caller = bool(
|
||||
run_config.get("external_api_caller") or getattr(self, "external_caller", False)
|
||||
)
|
||||
approval_required = bool(run_config.get("public_link_caller"))
|
||||
|
||||
wiki_config: Optional[Dict[str, Any]] = None
|
||||
try:
|
||||
@@ -1237,6 +1310,9 @@ class StreamProcessor:
|
||||
"source_owner_id": owner,
|
||||
"decoded_token": self.decoded_token,
|
||||
"user": caller,
|
||||
"outside_caller": outside_caller,
|
||||
"writes_allowed": not outside_caller or outside_edits_allowed(source_doc),
|
||||
"approval_required": approval_required,
|
||||
}
|
||||
break
|
||||
except Exception:
|
||||
@@ -1343,7 +1419,16 @@ class StreamProcessor:
|
||||
return None, None
|
||||
|
||||
def pre_fetch_tools(self) -> Optional[Dict[str, Any]]:
|
||||
"""Pre-fetch tool data for template rendering before agent creation"""
|
||||
"""Pre-fetch tool data for template rendering before agent creation.
|
||||
|
||||
Runs the actions the prompt template names on the toolset the agent
|
||||
run gets, so a teammate or public-link user renders the owner's
|
||||
prompt with the owner's tools, never their own.
|
||||
|
||||
Returns:
|
||||
Action results keyed by tool name and tool id, or None when
|
||||
nothing was fetched.
|
||||
"""
|
||||
if not settings.ENABLE_TOOL_PREFETCH:
|
||||
logger.info(
|
||||
"Tool pre-fetching disabled globally via ENABLE_TOOL_PREFETCH setting"
|
||||
@@ -1359,17 +1444,17 @@ class StreamProcessor:
|
||||
|
||||
try:
|
||||
user_id = self.initial_user_id or "local"
|
||||
agentless = self.agent_id is None
|
||||
with db_readonly() as conn:
|
||||
user_tools = UserToolsRepository(conn).list_active_for_user(user_id)
|
||||
user_doc = (
|
||||
UsersRepository(conn).get(user_id) if agentless else None
|
||||
)
|
||||
|
||||
default_docs = (
|
||||
synthesized_default_tools(user_doc) if agentless else []
|
||||
outside_caller = bool(
|
||||
self.agent_config.get("external_api_caller") or self.agent_config.get("public_link_caller")
|
||||
)
|
||||
tool_docs = list(user_tools) + default_docs
|
||||
# The same toolset the run gets: an agent's own tools (resolved as
|
||||
# its owner, or the editor who attached them), else the caller's
|
||||
# tools plus defaults. Explicit rows first, so they claim names.
|
||||
run_tools = [
|
||||
tool for tool in self._run_tool_executor().get_tools().values()
|
||||
if isinstance(tool, dict) and not tool.get("client_side")
|
||||
]
|
||||
tool_docs = sorted(run_tools, key=lambda tool: bool(tool.get("default")))
|
||||
if not tool_docs:
|
||||
return None
|
||||
|
||||
@@ -1397,6 +1482,15 @@ class StreamProcessor:
|
||||
continue
|
||||
required_actions = None
|
||||
|
||||
owner = tool_doc.get("user_id")
|
||||
if owner and (owner != user_id or outside_caller):
|
||||
# Someone else's tool (a widget or API run carries the
|
||||
# owner's id but isn't the owner): pre-fetch asks nobody,
|
||||
# so only what the run would do without asking.
|
||||
required_actions = self._unasked_actions(tool_doc, required_actions)
|
||||
if not required_actions:
|
||||
continue
|
||||
|
||||
tool_data = self._fetch_tool_data(tool_doc, required_actions)
|
||||
if tool_data:
|
||||
# Explicit rows claim the name key; a default tool takes
|
||||
@@ -1413,6 +1507,57 @@ class StreamProcessor:
|
||||
logger.warning(f"Failed to pre-fetch tools: {type(e).__name__}")
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _unasked_actions(
|
||||
tool_doc: Dict[str, Any], required_actions: Optional[Set[Optional[str]]]
|
||||
) -> Set[Optional[str]]:
|
||||
"""The required actions of someone else's tool that run without asking.
|
||||
|
||||
A tool on someone else's connected account runs on their account or
|
||||
needs the caller's own connection, a tool that decides approval per
|
||||
call (a remote device, the code executor) can't be judged from its
|
||||
stored flags, an approval-gated action waits for a person, and a
|
||||
write with the owner's credentials needs the owner's say-so;
|
||||
pre-fetch has none of these, so all are left out.
|
||||
|
||||
Args:
|
||||
tool_doc: The tool row, owned by someone other than the caller.
|
||||
required_actions: Action names the template needs; None, or a set
|
||||
holding None, means all of them.
|
||||
|
||||
Returns:
|
||||
The action names to run, empty when there are none.
|
||||
"""
|
||||
from docsgpt.connectors.permissions import owner_credential_writes, tool_actions
|
||||
|
||||
if tool_doc.get("connection_id") or tool_doc.get("name") in _LIVE_APPROVAL_TOOLS:
|
||||
return set()
|
||||
owner_writes = set(owner_credential_writes(tool_doc))
|
||||
unasked = {
|
||||
action.get("name") for action in tool_actions(tool_doc)
|
||||
if action.get("name") and action.get("active", True) and not action.get("require_approval")
|
||||
and action.get("name") not in owner_writes
|
||||
}
|
||||
if required_actions is None or None in required_actions:
|
||||
return unasked
|
||||
return {name for name in required_actions if name in unasked}
|
||||
|
||||
def _run_tool_executor(self):
|
||||
"""A ``ToolExecutor`` resolving the toolset this turn's agent run gets.
|
||||
|
||||
Returns:
|
||||
ToolExecutor: Built with the run's key, user and agent.
|
||||
"""
|
||||
from docsgpt.agents.tool_executor import ToolExecutor
|
||||
|
||||
user = self.decoded_token.get("sub") if self.decoded_token else None
|
||||
return ToolExecutor(
|
||||
user_api_key=self.agent_config.get("user_api_key"),
|
||||
user=user,
|
||||
decoded_token=self.decoded_token,
|
||||
agent_id=self.agent_id,
|
||||
)
|
||||
|
||||
def _enabled_tool_names(self) -> Optional[set]:
|
||||
"""Resolve the tool names enabled for this turn, for ``tools.enabled`` gating.
|
||||
|
||||
@@ -1422,15 +1567,7 @@ class StreamProcessor:
|
||||
(keeps the section) rather than hiding guidance when resolution breaks.
|
||||
"""
|
||||
try:
|
||||
from docsgpt.agents.tool_executor import ToolExecutor
|
||||
|
||||
user = self.decoded_token.get("sub") if self.decoded_token else None
|
||||
tool_executor = ToolExecutor(
|
||||
user_api_key=self.agent_config.get("user_api_key"),
|
||||
user=user,
|
||||
decoded_token=self.decoded_token,
|
||||
agent_id=self.agent_id,
|
||||
)
|
||||
tool_executor = self._run_tool_executor()
|
||||
client_tools = self.data.get("client_tools")
|
||||
if client_tools:
|
||||
tool_executor.client_tools = client_tools
|
||||
@@ -1638,21 +1775,37 @@ class StreamProcessor:
|
||||
from docsgpt.llm.handlers.handler_creator import LLMHandlerCreator
|
||||
from docsgpt.llm.llm_creator import LLMCreator
|
||||
|
||||
# Who is resuming, classified from this request alone: the saved state
|
||||
# says who paused the turn, but anyone holding the agent's key (a
|
||||
# widget key is public) can send the tool actions that resume it.
|
||||
request_key = self.data.get("api_key")
|
||||
original_token = self.decoded_token
|
||||
key_agent = None
|
||||
if request_key:
|
||||
with db_readonly() as conn:
|
||||
key_agent = AgentsRepository(conn).find_by_key(request_key)
|
||||
key_owner = (
|
||||
(key_agent.get("user_id") or key_agent.get("user")) if key_agent else None
|
||||
)
|
||||
request_external = bool(getattr(self, "external_caller", False)) or (
|
||||
bool(request_key) and is_external_api_caller(self.data, original_token, key_owner)
|
||||
)
|
||||
request_public_link = False
|
||||
named_agent = self.data.get("agent_id")
|
||||
if named_agent and not request_key:
|
||||
try:
|
||||
self._get_agent_key(str(named_agent), self.initial_user_id)
|
||||
except Exception as exc:
|
||||
raise ValueError("This conversation can't be resumed with that agent") from exc
|
||||
request_public_link = bool(getattr(self, "public_link_usage", False))
|
||||
|
||||
# api_key-in-body auth carries no JWT, so initial_user_id is None — but
|
||||
# the state was saved under the agent owner. Resolve the owner so the
|
||||
# lookup / mark_resuming / delete_state key on the same id. (No-op for
|
||||
# v1, which already passes an owner-scoped decoded_token.)
|
||||
if self.initial_user_id is None and self.data.get("api_key"):
|
||||
with db_readonly() as conn:
|
||||
agent_doc = AgentsRepository(conn).find_by_key(self.data["api_key"])
|
||||
owner = (
|
||||
(agent_doc.get("user_id") or agent_doc.get("user"))
|
||||
if agent_doc
|
||||
else None
|
||||
)
|
||||
if owner:
|
||||
self.initial_user_id = owner
|
||||
self.decoded_token = {"sub": owner}
|
||||
if self.initial_user_id is None and key_owner:
|
||||
self.initial_user_id = key_owner
|
||||
self.decoded_token = {"sub": key_owner}
|
||||
|
||||
cont_service = ContinuationService()
|
||||
state = claimed_state or cont_service.claim_state(
|
||||
@@ -1661,6 +1814,21 @@ class StreamProcessor:
|
||||
if not state:
|
||||
raise ValueError("No pending tool state found for this conversation")
|
||||
|
||||
# A request that names an agent (by key or id) resumes only that
|
||||
# agent's turn; the claim goes back so its rightful caller can resume.
|
||||
saved_agent = str((state.get("agent_config") or {}).get("agent_id") or "").lower()
|
||||
targets = []
|
||||
if request_key:
|
||||
targets.append(str((key_agent or {}).get("id") or (key_agent or {}).get("_id") or ""))
|
||||
if named_agent:
|
||||
targets.append(str(named_agent))
|
||||
if any(target.lower() != saved_agent or not target for target in targets):
|
||||
try:
|
||||
cont_service.release_claim(conversation_id, self.initial_user_id)
|
||||
except Exception:
|
||||
logger.warning("Failed to release a refused resume claim", exc_info=True)
|
||||
raise ValueError("This conversation belongs to a different agent")
|
||||
|
||||
messages = state["messages"]
|
||||
pending_tool_calls = state["pending_tool_calls"]
|
||||
tools_dict = state["tools_dict"]
|
||||
@@ -1696,11 +1864,20 @@ class StreamProcessor:
|
||||
if callable(importer):
|
||||
importer(agent_config.get("responses_state"))
|
||||
llm_handler = LLMHandlerCreator.create_handler(llm_name or "default")
|
||||
# Outside if either who paused the turn or who resumes it is.
|
||||
resume_external = bool(agent_config.get("external_api_caller")) or request_external
|
||||
resume_public_link = bool(agent_config.get("public_link_caller")) or request_public_link
|
||||
apply_resume_caller_rules(
|
||||
tools_dict, outside_caller=resume_external, public_link_caller=resume_public_link,
|
||||
)
|
||||
tool_executor = ToolExecutor(
|
||||
user_api_key=user_api_key,
|
||||
user=self.initial_user_id,
|
||||
decoded_token=self.decoded_token,
|
||||
agent_id=agent_id,
|
||||
external_caller=resume_external,
|
||||
public_link_caller=resume_public_link,
|
||||
api_write_allowlist=agent_config.get("api_write_allowlist"),
|
||||
)
|
||||
tool_executor.conversation_id = conversation_id
|
||||
# Restore client tools so they stay available for subsequent LLM calls
|
||||
@@ -1878,6 +2055,9 @@ class StreamProcessor:
|
||||
user=user,
|
||||
decoded_token=self.decoded_token,
|
||||
agent_id=self.agent_id,
|
||||
external_caller=bool(self.agent_config.get("external_api_caller")),
|
||||
public_link_caller=bool(self.agent_config.get("public_link_caller")),
|
||||
api_write_allowlist=self.agent_config.get("api_write_allowlist"),
|
||||
)
|
||||
tool_executor.conversation_id = self.conversation_id
|
||||
# Pass client-side tools so they get merged in get_tools()
|
||||
|
||||
@@ -0,0 +1,722 @@
|
||||
"""Connectors catalog and connections API.
|
||||
|
||||
``/api/connectors/catalog`` lists every service DocsGPT can connect to, with
|
||||
whether the server is set up for it and the caller's connection summary.
|
||||
``/api/connections`` lists and manages the caller's connections. Responses
|
||||
never include tokens or secrets.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from flask import current_app, jsonify, make_response, request
|
||||
from flask_restx import Namespace, Resource
|
||||
|
||||
import uuid
|
||||
|
||||
from docsgpt.api import api
|
||||
from docsgpt.api.user.authz import ROLE_ADMIN, has_role
|
||||
from docsgpt.connectors import catalog, service
|
||||
from docsgpt.storage.db.repositories.connector_sessions import ConnectorSessionsRepository
|
||||
from docsgpt.storage.db.session import db_readonly, db_session
|
||||
from docsgpt.security.encryption import CredentialDecryptionError
|
||||
|
||||
_FREQUENCIES = ("never", "daily", "weekly", "monthly")
|
||||
|
||||
connections_ns = Namespace("connections", description="Connectors and connections", path="/api")
|
||||
api.add_namespace(connections_ns)
|
||||
|
||||
|
||||
def _user_id() -> str | None:
|
||||
token = getattr(request, "decoded_token", None)
|
||||
return token.get("sub") if isinstance(token, dict) else None
|
||||
|
||||
|
||||
def _unauthorized():
|
||||
return make_response(jsonify({"success": False, "error": "Unauthorized"}), 401)
|
||||
|
||||
|
||||
def _not_found():
|
||||
return make_response(jsonify({"success": False, "error": "Connection not found"}), 404)
|
||||
|
||||
|
||||
@connections_ns.route("/connectors/catalog")
|
||||
class ConnectorCatalog(Resource):
|
||||
@api.doc(description="Every connector with its availability and the caller's connection summary")
|
||||
def get(self):
|
||||
user_id = _user_id()
|
||||
if not user_id:
|
||||
return _unauthorized()
|
||||
try:
|
||||
with db_readonly() as conn:
|
||||
entries = service.catalog_for_user(
|
||||
conn, user_id, is_admin=has_role(request.decoded_token, ROLE_ADMIN),
|
||||
)
|
||||
except Exception as err:
|
||||
current_app.logger.error(f"Error building connector catalog: {err}", exc_info=True)
|
||||
return make_response(jsonify({"success": False, "error": "Failed to load connectors"}), 500)
|
||||
return make_response(jsonify({"success": True, "connectors": entries}), 200)
|
||||
|
||||
|
||||
def _json_body() -> dict:
|
||||
body = request.get_json(silent=True)
|
||||
return body if isinstance(body, dict) else {}
|
||||
|
||||
|
||||
def _owned(conn, connection_id: str, user_id: str):
|
||||
return ConnectorSessionsRepository(conn).get_for_user(connection_id, user_id)
|
||||
|
||||
|
||||
def _error(message: str, status: int, **extra):
|
||||
return make_response(jsonify({"success": False, "error": message, **extra}), status)
|
||||
|
||||
|
||||
@connections_ns.route("/connections")
|
||||
class ConnectionsList(Resource):
|
||||
@api.doc(
|
||||
description=(
|
||||
"Create a connection from pasted credentials: "
|
||||
"{connector_key, credentials, label?}. Same credentials reuse the same connection."
|
||||
)
|
||||
)
|
||||
def post(self):
|
||||
user_id = _user_id()
|
||||
if not user_id:
|
||||
return _unauthorized()
|
||||
body = _json_body()
|
||||
definition = catalog.get_definition(body.get("connector_key"))
|
||||
if definition is None or definition.auth_kind != "api_key":
|
||||
return _error("This connector does not take pasted credentials", 400)
|
||||
if definition.missing_settings:
|
||||
return _error("This connector needs admin setup", 400, code="needs_setup")
|
||||
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=label,
|
||||
)
|
||||
except service.EncryptionKeyNotConfigured as err:
|
||||
return _error(str(err), 400, code="encryption_key_default")
|
||||
except service.ConnectorDisabled as err:
|
||||
return _error(str(err), 403, code="disabled")
|
||||
except ValueError as err:
|
||||
return _error(str(err), 400)
|
||||
except Exception as err:
|
||||
current_app.logger.error(f"Error creating connection: {err}", exc_info=True)
|
||||
return _error("Failed to create connection", 500)
|
||||
return make_response(
|
||||
jsonify(
|
||||
{
|
||||
"success": True,
|
||||
"created": created,
|
||||
"connection": service.serialize_connection(row),
|
||||
"setup": dict(definition.setup),
|
||||
}
|
||||
),
|
||||
201 if created else 200,
|
||||
)
|
||||
|
||||
@api.doc(description="The caller's connections with status and linked resource counts")
|
||||
def get(self):
|
||||
user_id = _user_id()
|
||||
if not user_id:
|
||||
return _unauthorized()
|
||||
try:
|
||||
with db_readonly() as conn:
|
||||
connections = service.list_connections(conn, user_id)
|
||||
except Exception as err:
|
||||
current_app.logger.error(f"Error listing connections: {err}", exc_info=True)
|
||||
return make_response(jsonify({"success": False, "error": "Failed to load connections"}), 500)
|
||||
return make_response(jsonify({"success": True, "connections": connections}), 200)
|
||||
|
||||
|
||||
@connections_ns.route("/connections/<string:connection_id>")
|
||||
class ConnectionDetail(Resource):
|
||||
@api.doc(description="One connection with the sources it syncs and the tools it provides")
|
||||
def get(self, connection_id: str):
|
||||
user_id = _user_id()
|
||||
if not user_id:
|
||||
return _unauthorized()
|
||||
try:
|
||||
with db_readonly() as conn:
|
||||
row = ConnectorSessionsRepository(conn).get_for_user(connection_id, user_id)
|
||||
if row is None:
|
||||
return _not_found()
|
||||
detail = service.connection_detail(conn, row)
|
||||
except Exception as err:
|
||||
current_app.logger.error(f"Error loading connection: {err}", exc_info=True)
|
||||
return make_response(jsonify({"success": False, "error": "Failed to load connection"}), 500)
|
||||
return make_response(jsonify({"success": True, "connection": detail}), 200)
|
||||
|
||||
@api.doc(description="Name an account: {name}. An empty name clears it. Owner only.")
|
||||
def patch(self, connection_id: str):
|
||||
user_id = _user_id()
|
||||
if not user_id:
|
||||
return _unauthorized()
|
||||
name = _json_body().get("name")
|
||||
if not isinstance(name, str) or len(name.strip()) > service.ACCOUNT_NAME_MAX:
|
||||
return _error(f"name must be text of at most {service.ACCOUNT_NAME_MAX} characters", 400)
|
||||
with db_session() as conn:
|
||||
row = _owned(conn, connection_id, user_id)
|
||||
if row is None:
|
||||
return _not_found()
|
||||
connection = service.rename_connection(conn, row, name)
|
||||
return make_response(jsonify({"success": True, "connection": connection}), 200)
|
||||
|
||||
@api.doc(
|
||||
description=(
|
||||
"Remove a connection: {sources: keep | delete, tools: delete | keep}. "
|
||||
"Kept sources keep their content and stop syncing."
|
||||
)
|
||||
)
|
||||
def delete(self, connection_id: str):
|
||||
from docsgpt.api.user.sources.routes import delete_source
|
||||
|
||||
user_id = _user_id()
|
||||
if not user_id:
|
||||
return _unauthorized()
|
||||
body = _json_body()
|
||||
sources_mode = body.get("sources", "keep")
|
||||
tools_mode = body.get("tools", "delete")
|
||||
if sources_mode not in ("keep", "delete") or tools_mode not in ("keep", "delete"):
|
||||
return _error("sources and tools must be keep or delete", 400)
|
||||
try:
|
||||
with db_session() as conn:
|
||||
row = _owned(conn, connection_id, user_id)
|
||||
if row is None:
|
||||
return _not_found()
|
||||
to_delete = service.remove_connection(conn, row, sources=sources_mode, tools=tools_mode)
|
||||
failed = [str(doc["id"]) for doc in to_delete if not delete_source(user_id, doc)]
|
||||
except Exception as err:
|
||||
current_app.logger.error(f"Error removing connection: {err}", exc_info=True)
|
||||
return _error("Failed to remove connection", 500)
|
||||
return make_response(jsonify({"success": True, "failed_sources": failed}), 200)
|
||||
|
||||
|
||||
@connections_ns.route("/connections/<string:connection_id>/disconnect")
|
||||
class ConnectionDisconnect(Resource):
|
||||
@api.doc(description="Delete a connection's stored credentials; its sources and tools stay")
|
||||
def post(self, connection_id: str):
|
||||
user_id = _user_id()
|
||||
if not user_id:
|
||||
return _unauthorized()
|
||||
try:
|
||||
with db_session() as conn:
|
||||
row = ConnectorSessionsRepository(conn).get_for_user(connection_id, user_id)
|
||||
if row is None:
|
||||
return _not_found()
|
||||
connection = service.disconnect(conn, row)
|
||||
except Exception as err:
|
||||
current_app.logger.error(f"Error disconnecting connection: {err}", exc_info=True)
|
||||
return make_response(jsonify({"success": False, "error": "Failed to disconnect"}), 500)
|
||||
return make_response(jsonify({"success": True, "connection": connection}), 200)
|
||||
|
||||
|
||||
@connections_ns.route("/connections/<string:connection_id>/setup")
|
||||
class ConnectionSetup(Resource):
|
||||
@api.doc(
|
||||
description=(
|
||||
"Apply the connect wizard's choices: {create_tools, allow_writes?, tool_permissions?, "
|
||||
"sync?: {items, frequency, name?, config?}}. allow_writes points GitHub's tool at its write "
|
||||
"endpoint. config is the synced source's retrieval settings, validated like an upload's. Honours an Idempotency-Key header for the sync."
|
||||
)
|
||||
)
|
||||
def post(self, connection_id: str):
|
||||
user_id = _user_id()
|
||||
if not user_id:
|
||||
return _unauthorized()
|
||||
body = _json_body()
|
||||
create_tools = body.get("create_tools", True)
|
||||
allow_writes = body.get("allow_writes", False)
|
||||
if not isinstance(allow_writes, bool):
|
||||
return _error("allow_writes must be true or false", 400)
|
||||
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)
|
||||
forbidden = bool(row) and allow_writes and not service.writes_allowed(
|
||||
service.load_policies(conn), catalog.connector_key_for_row(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 forbidden:
|
||||
return _error("Changes through this connector are turned off by an admin", 403,
|
||||
code="writes_forbidden")
|
||||
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, writes=allow_writes)
|
||||
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()
|
||||
tools = []
|
||||
if create_tools:
|
||||
tools = service.ensure_connection_tools(
|
||||
conn, user_id, row, permissions=body.get("tool_permissions") or None,
|
||||
mcp_actions=mcp_actions, mcp_writes=allow_writes,
|
||||
)
|
||||
account_parameters = service.connection_parameters(row)
|
||||
tool_payload = [service.serialize_tool(tool, account_parameters) for tool in tools]
|
||||
sources = []
|
||||
if body.get("sync"):
|
||||
started = _start_sync(user_id, row, body["sync"])
|
||||
if isinstance(started, tuple):
|
||||
return _error(*started)
|
||||
sources.append(started)
|
||||
except Exception as err:
|
||||
current_app.logger.error(f"Error setting up connection: {err}", exc_info=True)
|
||||
return _error("Failed to set up connection", 500)
|
||||
return make_response(jsonify({"success": True, "tools": tool_payload, "sources": sources}), 200)
|
||||
|
||||
|
||||
def _start_sync(user_id: str, row: dict, sync: dict):
|
||||
"""Queue the first ingest of a source synced from ``row``.
|
||||
|
||||
Args:
|
||||
user_id: The connection's owner.
|
||||
row: The connection row.
|
||||
sync: ``{items, frequency, name?, config?}``. ``config`` is the source's
|
||||
retrieval settings (a ``SourceConfig``), applied as an upload's are.
|
||||
|
||||
Returns:
|
||||
The source summary, or ``(message, status)`` on a bad request.
|
||||
"""
|
||||
from docsgpt.api.user.sources.upload import (
|
||||
_claim_task_or_get_cached,
|
||||
_derive_source_id,
|
||||
_parse_source_config,
|
||||
_read_idempotency_key,
|
||||
_scoped_idempotency_key,
|
||||
)
|
||||
from docsgpt.api.user.tasks import ingest_connector_task, ingest_remote
|
||||
|
||||
definition = catalog.get_definition(catalog.connector_key_for_row(row))
|
||||
if definition is None or not definition.sync_ingestor:
|
||||
return ("This connector does not sync content", 400)
|
||||
items = sync.get("items") or {}
|
||||
if not isinstance(items, dict):
|
||||
return ("items must be an object", 400)
|
||||
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()
|
||||
source_config, config_error = _parse_source_config(sync.get("config"))
|
||||
if config_error is not None:
|
||||
return ("Invalid source config", 400)
|
||||
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
|
||||
elif definition.sync_ingestor == "linear":
|
||||
from docsgpt.connectors import linear
|
||||
|
||||
try:
|
||||
items = linear.normalize_selection(items)
|
||||
except ValueError as err:
|
||||
return (str(err), 400)
|
||||
name = name or linear.selection_name(items)
|
||||
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":
|
||||
file_ids = [str(i) for i in items.get("file_ids") or [] if i]
|
||||
folder_ids = [str(i) for i in items.get("folder_ids") or [] if i]
|
||||
if not file_ids and not folder_ids:
|
||||
return ("Pick at least one file or folder", 400)
|
||||
task_fn = ingest_connector_task
|
||||
kwargs = {
|
||||
"job_name": name,
|
||||
"user": user_id,
|
||||
"source_type": definition.sync_ingestor,
|
||||
"connection_id": str(row["id"]),
|
||||
"file_ids": file_ids,
|
||||
"folder_ids": folder_ids,
|
||||
"recursive": bool(items.get("recursive", True)),
|
||||
"sync_frequency": frequency,
|
||||
"config": source_config,
|
||||
}
|
||||
else:
|
||||
if definition.sync_ingestor == "linear":
|
||||
# The teams and projects picked, read with the connection's MCP sign-in.
|
||||
source_data = items
|
||||
else:
|
||||
fields = {f.key for f in definition.setup_fields}
|
||||
source_data = {k: v for k, v in items.items() if k in fields and v not in (None, "")}
|
||||
missing = [f.label for f in definition.setup_fields if f.required and f.key not in source_data]
|
||||
if missing:
|
||||
return (f"Missing: {', '.join(missing)}", 400)
|
||||
task_fn = ingest_remote
|
||||
kwargs = {
|
||||
"source_data": source_data,
|
||||
"job_name": name,
|
||||
"user": user_id,
|
||||
"loader": definition.sync_ingestor,
|
||||
"connection_id": str(row["id"]),
|
||||
"sync_frequency": frequency,
|
||||
"config": source_config,
|
||||
}
|
||||
idempotency_key, _ = _read_idempotency_key()
|
||||
scoped_key = _scoped_idempotency_key(idempotency_key, user_id)
|
||||
task_id = None
|
||||
if scoped_key:
|
||||
task_id, cached = _claim_task_or_get_cached(scoped_key, "connection_setup_sync")
|
||||
if cached is not None:
|
||||
return {"id": cached.get("source_id"), "task_id": cached.get("task_id"), "name": name}
|
||||
source_id = str(_derive_source_id(scoped_key)) if scoped_key else str(uuid.uuid4())
|
||||
options = {"task_id": task_id} if task_id else {}
|
||||
task = task_fn.apply_async(
|
||||
kwargs={**kwargs, "idempotency_key": scoped_key, "source_id": source_id}, **options,
|
||||
)
|
||||
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>/linear")
|
||||
class LinearWorkspace(Resource):
|
||||
@api.doc(
|
||||
description=(
|
||||
"Linear: the teams and projects the connection can see, for the sync picker. "
|
||||
"Read through Linear's MCP server with the connection's sign-in."
|
||||
)
|
||||
)
|
||||
def get(self, connection_id: str):
|
||||
from docsgpt.connectors import linear, mcp
|
||||
|
||||
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) != linear.LINEAR_CONNECTOR:
|
||||
return _not_found()
|
||||
try:
|
||||
workspace = mcp.run_connection_session(row, linear.mcp_url(), linear.list_workspace)
|
||||
except service.ConnectionUnavailable:
|
||||
return _error("Reconnect to continue", 409, code="reconnect")
|
||||
except service.TransientConnectionError:
|
||||
return _error("Linear is not responding. Try again.", 503)
|
||||
except Exception as err:
|
||||
current_app.logger.error(f"Error listing Linear teams: {err}", exc_info=True)
|
||||
return _error("Failed to list Linear teams", 502)
|
||||
return make_response(jsonify({"success": True, **workspace}), 200)
|
||||
|
||||
|
||||
@connections_ns.route("/connections/<string:connection_id>/reconnect")
|
||||
class ConnectionReconnect(Resource):
|
||||
@api.doc(
|
||||
description=(
|
||||
"OAuth: returns an authorization URL for the same account. "
|
||||
"API key: accepts {credentials} and replaces the stored ones."
|
||||
)
|
||||
)
|
||||
def post(self, connection_id: str):
|
||||
from docsgpt.api.connector.routes import build_authorization
|
||||
|
||||
user_id = _user_id()
|
||||
if not user_id:
|
||||
return _unauthorized()
|
||||
body = _json_body()
|
||||
try:
|
||||
with db_readonly() as conn:
|
||||
row = _owned(conn, connection_id, user_id)
|
||||
if row is None:
|
||||
return _not_found()
|
||||
key = catalog.connector_key_for_row(row)
|
||||
definition = catalog.get_definition(key)
|
||||
auth_kind = row.get("auth_kind") or (definition.auth_kind if definition else None)
|
||||
if auth_kind == "oauth":
|
||||
started = build_authorization(row["provider"], user_id, connection_id)
|
||||
return make_response(jsonify({"success": True, "kind": "oauth", **started}), 200)
|
||||
if auth_kind == "mcp_oauth":
|
||||
# The MCP client runs the OAuth dance (dynamic registration,
|
||||
# PKCE); the frontend starts it through /api/mcp_server/test.
|
||||
return make_response(
|
||||
jsonify({"success": True, "kind": "mcp_oauth", "server_url": row.get("server_url")}), 200,
|
||||
)
|
||||
credentials = body.get("credentials")
|
||||
if not isinstance(credentials, dict) or not credentials:
|
||||
return _error("credentials are required", 400)
|
||||
service.ensure_can_store_credentials()
|
||||
with db_session() as conn:
|
||||
locked = ConnectorSessionsRepository(conn).get_for_update(connection_id)
|
||||
# read_secrets, not load_secrets: flagging an unreadable row
|
||||
# would write it from a second transaction while this one
|
||||
# holds its lock. The new credentials replace it anyway.
|
||||
try:
|
||||
stored = service.read_secrets(locked)
|
||||
except CredentialDecryptionError:
|
||||
stored = {}
|
||||
merged = {**(stored.get("credentials") or {}), **{k: v for k, v in credentials.items() if v}}
|
||||
service.write_secrets(
|
||||
conn, locked, {**stored, "credentials": merged},
|
||||
status=service.STATUS_CONNECTED, last_error=None,
|
||||
)
|
||||
service.resume_sources(conn, connection_id)
|
||||
connection = service.serialize_connection(ConnectorSessionsRepository(conn).get(connection_id))
|
||||
except service.EncryptionKeyNotConfigured as err:
|
||||
return _error(str(err), 400, code="encryption_key_default")
|
||||
except Exception as err:
|
||||
current_app.logger.error(f"Error reconnecting: {err}", exc_info=True)
|
||||
return _error("Failed to reconnect", 500)
|
||||
return make_response(jsonify({"success": True, "kind": "api_key", "connection": connection}), 200)
|
||||
|
||||
|
||||
@connections_ns.route("/connections/<string:connection_id>/picker-token")
|
||||
class ConnectionPickerToken(Resource):
|
||||
@api.doc(description="A short-lived access token for a browser-side file picker. Owner only; never a refresh token.")
|
||||
def post(self, connection_id: str):
|
||||
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 (row.get("auth_kind") or "oauth") != "oauth":
|
||||
return _not_found()
|
||||
try:
|
||||
token = service.picker_token(connection_id)
|
||||
except service.ConnectionUnavailable:
|
||||
return _error("Reconnect to continue", 409, code="reconnect")
|
||||
except service.TransientConnectionError:
|
||||
return _error("The provider is not responding. Try again.", 503)
|
||||
return make_response(jsonify({"success": True, **token}), 200)
|
||||
|
||||
|
||||
@connections_ns.route("/connections/claim")
|
||||
class ConnectionClaim(Resource):
|
||||
@api.doc(
|
||||
description=(
|
||||
"One-time link of a legacy browser session token ({provider, session_token}) "
|
||||
"to the caller's connection. Removed next release."
|
||||
)
|
||||
)
|
||||
def post(self):
|
||||
user_id = _user_id()
|
||||
if not user_id:
|
||||
return _unauthorized()
|
||||
body = _json_body()
|
||||
provider, token = body.get("provider"), body.get("session_token")
|
||||
if not provider or not token:
|
||||
return _error("provider and session_token are required", 400)
|
||||
with db_readonly() as conn:
|
||||
row = service.claim_session_token(conn, user_id, str(provider), str(token))
|
||||
if row is None:
|
||||
return _not_found()
|
||||
return make_response(jsonify({"success": True, "connection_id": str(row["id"])}), 200)
|
||||
|
||||
|
||||
@connections_ns.route("/connections/<string:connection_id>/tools/<string:tool_id>/permissions")
|
||||
class ConnectionToolPermissions(Resource):
|
||||
@api.doc(description="Set per-action permissions: {permissions: {action: always | ask | off}}")
|
||||
def put(self, connection_id: str, tool_id: str):
|
||||
from docsgpt.connectors.permissions import PERMISSIONS
|
||||
|
||||
user_id = _user_id()
|
||||
if not user_id:
|
||||
return _unauthorized()
|
||||
permissions = _json_body().get("permissions")
|
||||
if not isinstance(permissions, dict) or any(p not in PERMISSIONS for p in permissions.values()):
|
||||
return _error("permissions must map action names to always, ask or off", 400)
|
||||
with db_session() as conn:
|
||||
row = _owned(conn, connection_id, user_id)
|
||||
if row is None:
|
||||
return _not_found()
|
||||
tool = service.set_tool_permissions(conn, user_id, connection_id, tool_id, permissions)
|
||||
if tool is None:
|
||||
return _not_found()
|
||||
payload = service.serialize_tool(tool, service.connection_parameters(row))
|
||||
return make_response(jsonify({"success": True, "tool": payload}), 200)
|
||||
|
||||
|
||||
@connections_ns.route("/connections/<string:connection_id>/tools/<string:tool_id>/parameters")
|
||||
class ConnectionToolParameters(Resource):
|
||||
@api.doc(
|
||||
description=(
|
||||
"Fix or release an action's parameters: {action, parameters: {name: value | null}}. "
|
||||
"A value is sent on every call and hidden from the model; null lets the model decide. Owner only."
|
||||
)
|
||||
)
|
||||
def put(self, connection_id: str, tool_id: str):
|
||||
user_id = _user_id()
|
||||
if not user_id:
|
||||
return _unauthorized()
|
||||
body = _json_body()
|
||||
action, pins = body.get("action"), body.get("parameters")
|
||||
if not isinstance(action, str) or not isinstance(pins, dict) or not pins:
|
||||
return _error("Send the action and a map of parameters to values or null", 400)
|
||||
with db_session() as conn:
|
||||
row = _owned(conn, connection_id, user_id)
|
||||
if row is None:
|
||||
return _not_found()
|
||||
try:
|
||||
tool = service.set_tool_parameters(conn, user_id, connection_id, tool_id, action, pins)
|
||||
except ValueError as err:
|
||||
return _error(str(err), 400)
|
||||
if tool is None:
|
||||
return _not_found()
|
||||
payload = service.serialize_tool(tool, service.connection_parameters(row))
|
||||
return make_response(jsonify({"success": True, "tool": payload}), 200)
|
||||
|
||||
|
||||
@connections_ns.route("/connections/<string:connection_id>/refresh-tools")
|
||||
class ConnectionRefreshTools(Resource):
|
||||
@api.doc(description="MCP: re-scan the server's actions and return what was added and removed")
|
||||
def post(self, connection_id: str):
|
||||
from docsgpt.connectors.mcp import refresh_mcp_tools
|
||||
|
||||
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:
|
||||
return _not_found()
|
||||
try:
|
||||
diff = refresh_mcp_tools(user_id, row)
|
||||
except service.ConnectionUnavailable:
|
||||
return _error("Reconnect to continue", 409, code="reconnect")
|
||||
except Exception as err:
|
||||
current_app.logger.error(f"Error refreshing MCP tools: {err}", exc_info=True)
|
||||
return _error("Failed to refresh tools", 502)
|
||||
return make_response(jsonify({"success": True, **diff}), 200)
|
||||
|
||||
|
||||
@connections_ns.route("/connections/<string:connection_id>/writes")
|
||||
class ConnectionWrites(Resource):
|
||||
@api.doc(
|
||||
description=(
|
||||
"GitHub: let agents make changes through this connection, or only read: {allow}. "
|
||||
"Re-reads the tool's actions from the matching endpoint. Owner only."
|
||||
)
|
||||
)
|
||||
def put(self, connection_id: str):
|
||||
from docsgpt.connectors.mcp import NoMcpTool, set_builtin_writes
|
||||
|
||||
user_id = _user_id()
|
||||
if not user_id:
|
||||
return _unauthorized()
|
||||
allow = _json_body().get("allow")
|
||||
if not isinstance(allow, bool):
|
||||
return _error("allow must be true or false", 400)
|
||||
with db_readonly() as conn:
|
||||
row = _owned(conn, connection_id, user_id)
|
||||
if row is None:
|
||||
return _not_found()
|
||||
try:
|
||||
result = set_builtin_writes(user_id, row, allow)
|
||||
except service.WritesForbidden as err:
|
||||
return _error(str(err), 403, code="writes_forbidden")
|
||||
except NoMcpTool as err:
|
||||
return _error(str(err), 409, code="no_tools")
|
||||
except ValueError as err:
|
||||
return _error(str(err), 400)
|
||||
except service.ConnectionUnavailable:
|
||||
return _error("Reconnect to continue", 409, code="reconnect")
|
||||
except Exception as err:
|
||||
current_app.logger.error(f"Error switching write access: {err}", exc_info=True)
|
||||
return _error("The service's tools could not be reached. Try again.", 502, code="tools_unavailable")
|
||||
return make_response(jsonify({"success": True, **result}), 200)
|
||||
|
||||
|
||||
@connections_ns.route("/connections/tools/<string:tool_id>/credential-mode")
|
||||
class ToolCredentialMode(Resource):
|
||||
@api.doc(
|
||||
description=(
|
||||
"Whose account a shared connection-backed tool uses: {mode: owner | member}. "
|
||||
"Owner only; refused when an admin forces a mode for the connector."
|
||||
)
|
||||
)
|
||||
def put(self, tool_id: str):
|
||||
from docsgpt.connectors.resolve import MODE_MEMBER, MODE_OWNER
|
||||
from docsgpt.storage.db.repositories.user_tools import UserToolsRepository
|
||||
|
||||
user_id = _user_id()
|
||||
if not user_id:
|
||||
return _unauthorized()
|
||||
mode = _json_body().get("mode")
|
||||
if mode not in (MODE_OWNER, MODE_MEMBER):
|
||||
return _error("mode must be owner or member", 400)
|
||||
with db_session() as conn:
|
||||
tools = UserToolsRepository(conn)
|
||||
tool = tools.get_any(tool_id, user_id)
|
||||
if tool is None or tool.get("user_id") != user_id or not tool.get("connection_id"):
|
||||
return _error("Tool not found", 404)
|
||||
connection = ConnectorSessionsRepository(conn).get(str(tool["connection_id"]))
|
||||
forced = service.forced_credential_mode(
|
||||
conn, catalog.connector_key_for_row(connection) if connection else None,
|
||||
)
|
||||
if forced and forced != mode:
|
||||
return _error("An admin sets this for every share", 409, code="forced", mode=forced)
|
||||
tools.update(str(tool["id"]), user_id, {"credential_mode": mode})
|
||||
return make_response(jsonify({"success": True, "mode": mode}), 200)
|
||||
+187
-143
@@ -1,7 +1,6 @@
|
||||
import base64
|
||||
import html
|
||||
import json
|
||||
import uuid
|
||||
from typing import Optional
|
||||
from urllib.parse import urlencode, urlsplit
|
||||
|
||||
@@ -22,6 +21,7 @@ from docsgpt.api.user.sources.access import load_source
|
||||
from docsgpt.api.user.tasks import (
|
||||
ingest_connector_task,
|
||||
)
|
||||
from docsgpt.connectors import service
|
||||
from docsgpt.core.settings import settings
|
||||
from docsgpt.parser.connectors.connector_creator import ConnectorCreator
|
||||
from docsgpt.storage.db.repositories.connector_sessions import (
|
||||
@@ -101,17 +101,25 @@ def _js_literal(value) -> str:
|
||||
|
||||
def _render_callback_page(
|
||||
status: str, message: str, provider_raw: str, session_token: str = "", user_email: str = "",
|
||||
connection_id: str = "",
|
||||
):
|
||||
"""Popup page that reports an OAuth result to the opener on allowed origins only."""
|
||||
status = status if status in ("success", "error", "cancelled") else "error"
|
||||
# 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:
|
||||
payload = {"type": f"{provider_key}_auth_success", "session_token": session_token, "user_email": user_email}
|
||||
payload = {
|
||||
"type": f"{provider_key}_auth_success",
|
||||
# The connection id is what current frontends use; the session
|
||||
# token is kept for one release for frontends from before it.
|
||||
"connection_id": connection_id,
|
||||
"session_token": session_token,
|
||||
"user_email": user_email,
|
||||
}
|
||||
elif provider_key and status == "error":
|
||||
# The frontend shows its own localized failure message; cancellations are
|
||||
# reported when the popup closes.
|
||||
@@ -171,6 +179,42 @@ def _render_callback_page(
|
||||
|
||||
|
||||
|
||||
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:
|
||||
session_row = service.begin_oauth(conn, user_id, provider, connection_id)
|
||||
state = base64.urlsafe_b64encode(
|
||||
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": url,
|
||||
"state": state,
|
||||
"callback_origin": _origin_of(settings.CONNECTOR_REDIRECT_BASE_URI),
|
||||
}
|
||||
|
||||
|
||||
@connectors_ns.route("/api/connectors/auth")
|
||||
class ConnectorAuth(Resource):
|
||||
@api.doc(description="Get connector OAuth authorization URL", params={"provider": "Connector provider (e.g., google_drive)"})
|
||||
@@ -180,27 +224,26 @@ 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
|
||||
if not decoded_token:
|
||||
return make_response(jsonify({"success": False, "error": "Unauthorized"}), 401)
|
||||
user_id = decoded_token.get('sub')
|
||||
|
||||
with db_session() as conn:
|
||||
session_row = ConnectorSessionsRepository(conn).upsert(
|
||||
user_id, provider, status="pending",
|
||||
try:
|
||||
started = build_authorization(
|
||||
provider, user_id, request.args.get("connection_id") or None,
|
||||
install=request.args.get("install") in ("1", "true"),
|
||||
)
|
||||
session_pg_id = str(session_row["id"])
|
||||
state_dict = {
|
||||
"provider": provider,
|
||||
"object_id": session_pg_id,
|
||||
}
|
||||
state = base64.urlsafe_b64encode(json.dumps(state_dict).encode()).decode()
|
||||
|
||||
auth = ConnectorCreator.create_auth(provider)
|
||||
authorization_url = auth.get_authorization_url(state=state)
|
||||
except service.EncryptionKeyNotConfigured as err:
|
||||
return make_response(
|
||||
jsonify({"success": False, "error": str(err), "code": "encryption_key_default"}), 400,
|
||||
)
|
||||
except service.ConnectionUnavailable:
|
||||
return make_response(jsonify({"success": False, "error": "Connection not found"}), 404)
|
||||
except service.ConnectorDisabled as err:
|
||||
return make_response(jsonify({"success": False, "error": str(err), "code": "disabled"}), 403)
|
||||
# The popup drops results for origins outside the allowlist, which the
|
||||
# user only sees as a cancelled sign-in; name the missing origin here.
|
||||
request_origin = _origin_of(request.headers.get("Origin"))
|
||||
@@ -209,12 +252,7 @@ class ConnectorAuth(Resource):
|
||||
f"Connector sign-in requested from {request_origin}, which cannot receive the result; "
|
||||
"add it to CONNECTOR_ALLOWED_ORIGINS"
|
||||
)
|
||||
return make_response(jsonify({
|
||||
"success": True,
|
||||
"authorization_url": authorization_url,
|
||||
"state": state,
|
||||
"callback_origin": _origin_of(settings.CONNECTOR_REDIRECT_BASE_URI),
|
||||
}), 200)
|
||||
return make_response(jsonify({"success": True, **started}), 200)
|
||||
except Exception as e:
|
||||
current_app.logger.error(f"Error generating connector auth URL: {e}", exc_info=True)
|
||||
return make_response(jsonify({"success": False, "error": "Failed to generate authorization URL"}), 500)
|
||||
@@ -233,12 +271,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"
|
||||
@@ -270,16 +320,16 @@ class ConnectorsCallback(Resource):
|
||||
auth = ConnectorCreator.create_auth(provider)
|
||||
token_info = auth.exchange_code_for_tokens(authorization_code)
|
||||
|
||||
session_token = str(uuid.uuid4())
|
||||
|
||||
try:
|
||||
if provider == "google_drive":
|
||||
credentials = auth.create_credentials_from_token_info(token_info)
|
||||
service = auth.build_drive_service(credentials)
|
||||
user_info = service.about().get(fields="user").execute()
|
||||
drive_service = auth.build_drive_service(credentials)
|
||||
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}")
|
||||
@@ -289,29 +339,23 @@ class ConnectorsCallback(Resource):
|
||||
|
||||
# ``object_id`` in the OAuth state is the PG session row
|
||||
# UUID (new flow) or a legacy Mongo ObjectId (pre-cutover
|
||||
# issued state). Try UUID update first; fall back to
|
||||
# legacy id path.
|
||||
patch = {
|
||||
"session_token": session_token,
|
||||
"token_info": sanitized_token_info,
|
||||
"user_email": user_email,
|
||||
"status": "authorized",
|
||||
}
|
||||
# issued state).
|
||||
with db_session() as conn:
|
||||
repo = ConnectorSessionsRepository(conn)
|
||||
if state_object_id:
|
||||
value = str(state_object_id)
|
||||
updated = False
|
||||
if len(value) == 36 and "-" in value:
|
||||
updated = repo.update(value, patch)
|
||||
if not updated:
|
||||
repo.update_by_legacy_id(value, patch)
|
||||
value = str(state_object_id or "")
|
||||
state_row = repo.get(value) or (repo.get_by_legacy_id(value) if value else None)
|
||||
if state_row is None or state_row.get("provider") != provider:
|
||||
raise ValueError("OAuth state names no pending connection")
|
||||
connection = service.complete_oauth(
|
||||
conn, state_row, provider, sanitized_token_info, user_email,
|
||||
)
|
||||
|
||||
# Render instead of redirecting so the session token never
|
||||
# lands in a URL (browser history, access logs, Referer).
|
||||
return _render_callback_page(
|
||||
"success", "Authentication successful", provider,
|
||||
session_token=session_token, user_email=user_email,
|
||||
session_token=connection.get("session_token") or "", user_email=user_email,
|
||||
connection_id=str(connection["id"]),
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
@@ -334,7 +378,8 @@ class ConnectorsCallback(Resource):
|
||||
class ConnectorFiles(Resource):
|
||||
@api.expect(api.model("ConnectorFilesModel", {
|
||||
"provider": fields.String(required=True),
|
||||
"session_token": fields.String(required=True),
|
||||
"connection_id": fields.String(required=False),
|
||||
"session_token": fields.String(required=False, description="Legacy; use connection_id"),
|
||||
"folder_id": fields.String(required=False),
|
||||
"limit": fields.Integer(required=False),
|
||||
"page_token": fields.String(required=False),
|
||||
@@ -345,26 +390,29 @@ class ConnectorFiles(Resource):
|
||||
try:
|
||||
data = request.get_json()
|
||||
provider = data.get('provider')
|
||||
session_token = data.get('session_token')
|
||||
limit = data.get('limit', 10)
|
||||
|
||||
if not provider or not session_token:
|
||||
return make_response(jsonify({"success": False, "error": "provider and session_token are required"}), 400)
|
||||
if not provider or not (data.get('connection_id') or data.get('session_token')):
|
||||
return make_response(
|
||||
jsonify({"success": False, "error": "provider and connection_id are required"}), 400,
|
||||
)
|
||||
|
||||
decoded_token = request.decoded_token
|
||||
if not decoded_token:
|
||||
return make_response(jsonify({"success": False, "error": "Unauthorized"}), 401)
|
||||
user = decoded_token.get('sub')
|
||||
with db_readonly() as conn:
|
||||
session = ConnectorSessionsRepository(conn).get_by_session_token(
|
||||
session_token,
|
||||
)
|
||||
if not owns_connector_session(session, user, provider):
|
||||
session = service.resolve_request_connection(user, provider, data)
|
||||
if session is None:
|
||||
return make_response(jsonify({"success": False, "error": "Invalid or unauthorized session"}), 401)
|
||||
|
||||
loader = ConnectorCreator.create_connector(provider, session_token)
|
||||
try:
|
||||
loader = ConnectorCreator.create_connector(provider, connection_id=str(session["id"]))
|
||||
except service.ConnectionUnavailable:
|
||||
return make_response(
|
||||
jsonify({"success": False, "error": "Reconnect to continue", "reconnect": True}), 401,
|
||||
)
|
||||
|
||||
generic_keys = {'provider', 'session_token'}
|
||||
generic_keys = {'provider', 'session_token', 'connection_id'}
|
||||
input_config = {
|
||||
k: v for k, v in data.items() if k not in generic_keys
|
||||
}
|
||||
@@ -409,65 +457,49 @@ class ConnectorFiles(Resource):
|
||||
|
||||
@connectors_ns.route("/api/connectors/validate-session")
|
||||
class ConnectorValidateSession(Resource):
|
||||
@api.expect(api.model("ConnectorValidateSessionModel", {"provider": fields.String(required=True), "session_token": fields.String(required=True)}))
|
||||
@api.doc(description="Validate connector session token and return user info and access token")
|
||||
@api.expect(api.model("ConnectorValidateSessionModel", {
|
||||
"provider": fields.String(required=True),
|
||||
"connection_id": fields.String(required=False),
|
||||
"session_token": fields.String(required=False, description="Legacy; use connection_id"),
|
||||
}))
|
||||
@api.doc(description="Validate a connection and return the account and a short-lived access token")
|
||||
def post(self):
|
||||
try:
|
||||
data = request.get_json()
|
||||
data = request.get_json() or {}
|
||||
provider = data.get('provider')
|
||||
session_token = data.get('session_token')
|
||||
if not provider or not session_token:
|
||||
return make_response(jsonify({"success": False, "error": "provider and session_token are required"}), 400)
|
||||
if not provider or not (data.get('connection_id') or data.get('session_token')):
|
||||
return make_response(
|
||||
jsonify({"success": False, "error": "provider and connection_id are required"}), 400,
|
||||
)
|
||||
|
||||
decoded_token = request.decoded_token
|
||||
if not decoded_token:
|
||||
return make_response(jsonify({"success": False, "error": "Unauthorized"}), 401)
|
||||
user = decoded_token.get('sub')
|
||||
|
||||
with db_readonly() as conn:
|
||||
session = ConnectorSessionsRepository(conn).get_by_session_token(
|
||||
session_token,
|
||||
)
|
||||
if not owns_connector_session(session, user, provider) or not session.get("token_info"):
|
||||
session = service.resolve_request_connection(user, provider, data)
|
||||
if session is None:
|
||||
return make_response(jsonify({"success": False, "error": "Invalid or expired session"}), 401)
|
||||
|
||||
token_info = session["token_info"]
|
||||
auth = ConnectorCreator.create_auth(provider)
|
||||
is_expired = auth.is_token_expired(token_info)
|
||||
|
||||
if is_expired and token_info.get('refresh_token'):
|
||||
try:
|
||||
refreshed_token_info = auth.refresh_access_token(token_info.get('refresh_token'))
|
||||
sanitized_token_info = auth.sanitize_token_info(refreshed_token_info)
|
||||
with db_session() as conn:
|
||||
repo = ConnectorSessionsRepository(conn)
|
||||
row = repo.get_by_session_token(session_token)
|
||||
if row:
|
||||
repo.update(str(row["id"]), {"token_info": sanitized_token_info})
|
||||
token_info = sanitized_token_info
|
||||
is_expired = False
|
||||
except Exception as refresh_error:
|
||||
current_app.logger.error(f"Failed to refresh token: {refresh_error}")
|
||||
|
||||
if is_expired:
|
||||
try:
|
||||
token = service.picker_token(str(session["id"]))
|
||||
except service.ConnectionUnavailable:
|
||||
return make_response(jsonify({
|
||||
"success": False,
|
||||
"expired": True,
|
||||
"error": "Session token has expired. Please reconnect."
|
||||
}), 401)
|
||||
except service.TransientConnectionError:
|
||||
return make_response(
|
||||
jsonify({"success": False, "error": "The provider is not responding. Try again."}), 503,
|
||||
)
|
||||
|
||||
_base_fields = {"access_token", "refresh_token", "token_uri", "expiry"}
|
||||
provider_extras = {k: v for k, v in token_info.items() if k not in _base_fields}
|
||||
|
||||
response_data = {
|
||||
return make_response(jsonify({
|
||||
"success": True,
|
||||
"expired": False,
|
||||
"user_email": session.get('user_email', 'Connected User'),
|
||||
"access_token": token_info.get('access_token'),
|
||||
**provider_extras,
|
||||
}
|
||||
|
||||
return make_response(jsonify(response_data), 200)
|
||||
"connection_id": str(session["id"]),
|
||||
"user_email": session.get('account_label') or session.get('user_email') or 'Connected User',
|
||||
**token,
|
||||
}), 200)
|
||||
except Exception as e:
|
||||
current_app.logger.error(f"Error validating connector session: {e}", exc_info=True)
|
||||
return make_response(jsonify({"success": False, "error": "Failed to validate session"}), 500)
|
||||
@@ -484,15 +516,13 @@ class ConnectorDisconnect(Resource):
|
||||
try:
|
||||
data = request.get_json()
|
||||
provider = data.get('provider')
|
||||
session_token = data.get('session_token')
|
||||
if not provider:
|
||||
return make_response(jsonify({"success": False, "error": "provider is required"}), 400)
|
||||
|
||||
if session_token:
|
||||
session = service.resolve_request_connection(decoded_token.get('sub'), provider, data)
|
||||
if session is not None:
|
||||
with db_session() as conn:
|
||||
ConnectorSessionsRepository(conn).delete_by_session_token(
|
||||
session_token, decoded_token.get('sub'),
|
||||
)
|
||||
service.disconnect(conn, session)
|
||||
|
||||
return make_response(jsonify({"success": True}), 200)
|
||||
except Exception as e:
|
||||
@@ -500,38 +530,51 @@ class ConnectorDisconnect(Resource):
|
||||
return make_response(jsonify({"success": False, "error": "Failed to disconnect session"}), 500)
|
||||
|
||||
|
||||
def _owner_connector_session(conn, owner_id: str, provider: str) -> Optional[dict]:
|
||||
"""The owner's usable connector session for ``provider``, or None.
|
||||
def _owner_connector_session(
|
||||
conn, owner_id: str, provider: str, connection_id: Optional[str] = None,
|
||||
) -> Optional[dict]:
|
||||
"""The owner's usable connection for ``provider``, or None.
|
||||
|
||||
Used when a team editor syncs a shared connector source: the sync runs
|
||||
with the owner's account. A session with no token, no stored credentials,
|
||||
or an expired access token that can't be refreshed counts as missing.
|
||||
with the owner's account, the source's own connection when it has one,
|
||||
else the owner's first connection for the provider. A connection that is
|
||||
not connected (signed out, flagged for reconnect, holding no
|
||||
credentials), cannot be decrypted, or has an expired access token that
|
||||
can't be refreshed counts as missing.
|
||||
|
||||
Args:
|
||||
conn: Open database connection.
|
||||
owner_id: The source owner's ``sub``.
|
||||
provider: The source's connector provider.
|
||||
connection_id: The source's own connection, if it names one.
|
||||
|
||||
Returns:
|
||||
Optional[dict]: The session row, or None when the owner must reconnect.
|
||||
Optional[dict]: The connection row, or None when the owner must reconnect.
|
||||
"""
|
||||
candidates = [
|
||||
s for s in ConnectorSessionsRepository(conn).list_for_user(owner_id)
|
||||
if owns_connector_session(s, owner_id, provider)
|
||||
and s.get("session_token") and s.get("token_info")
|
||||
]
|
||||
if not candidates:
|
||||
return None
|
||||
session = candidates[0]
|
||||
token_info = session["token_info"]
|
||||
if not token_info.get("refresh_token"):
|
||||
from docsgpt.security.encryption import CredentialDecryptionError
|
||||
|
||||
repo = ConnectorSessionsRepository(conn)
|
||||
if connection_id:
|
||||
row = repo.get_for_user(str(connection_id), owner_id)
|
||||
candidates = [row] if owns_connector_session(row, owner_id, provider) else []
|
||||
else:
|
||||
candidates = [s for s in repo.list_for_user(owner_id) if owns_connector_session(s, owner_id, provider)]
|
||||
for session in candidates:
|
||||
if service.normalize_status(session) != service.STATUS_CONNECTED:
|
||||
continue
|
||||
try:
|
||||
if ConnectorCreator.create_auth(provider).is_token_expired(token_info):
|
||||
return None
|
||||
except Exception:
|
||||
# Providers without an expiry check leave the verdict to the sync.
|
||||
pass
|
||||
return session
|
||||
token_info = service.read_secrets(session).get("token_info") or {}
|
||||
except CredentialDecryptionError:
|
||||
continue
|
||||
if isinstance(token_info, dict) and token_info and not token_info.get("refresh_token"):
|
||||
try:
|
||||
if ConnectorCreator.create_auth(provider).is_token_expired(token_info):
|
||||
continue
|
||||
except Exception:
|
||||
# Providers without an expiry check leave the verdict to the sync.
|
||||
pass
|
||||
return session
|
||||
return None
|
||||
|
||||
|
||||
@connectors_ns.route("/api/connectors/sync")
|
||||
@@ -541,11 +584,12 @@ class ConnectorSync(Resource):
|
||||
"ConnectorSyncModel",
|
||||
{
|
||||
"source_id": fields.String(required=True, description="Source ID to sync"),
|
||||
"session_token": fields.String(
|
||||
"connection_id": fields.String(
|
||||
required=False,
|
||||
description="The owner's connector session token (ignored for team editors, "
|
||||
"whose sync uses the owner's session)",
|
||||
)
|
||||
description="Connection to sync with; defaults to the source's own (ignored for "
|
||||
"team editors, whose sync uses the owner's connection)",
|
||||
),
|
||||
"session_token": fields.String(required=False, description="Legacy; use connection_id")
|
||||
},
|
||||
)
|
||||
)
|
||||
@@ -558,13 +602,12 @@ class ConnectorSync(Resource):
|
||||
try:
|
||||
data = request.get_json() or {}
|
||||
source_id = data.get('source_id')
|
||||
session_token = data.get('session_token')
|
||||
|
||||
if not source_id:
|
||||
return make_response(
|
||||
jsonify({
|
||||
"success": False,
|
||||
"error": "source_id and session_token are required"
|
||||
"error": "source_id is required"
|
||||
}),
|
||||
400
|
||||
)
|
||||
@@ -582,14 +625,6 @@ class ConnectorSync(Resource):
|
||||
)
|
||||
owner_id = ra.owner_id
|
||||
is_owner = ra.access == "owner"
|
||||
if is_owner and not session_token:
|
||||
return make_response(
|
||||
jsonify({
|
||||
"success": False,
|
||||
"error": "source_id and session_token are required"
|
||||
}),
|
||||
400
|
||||
)
|
||||
|
||||
remote_data = source.get('remote_data') or {}
|
||||
if isinstance(remote_data, str):
|
||||
@@ -609,17 +644,27 @@ class ConnectorSync(Resource):
|
||||
400
|
||||
)
|
||||
|
||||
source_connection = str(source['connection_id']) if source.get('connection_id') else None
|
||||
if is_owner:
|
||||
with db_readonly() as conn:
|
||||
session = ConnectorSessionsRepository(conn).get_by_session_token(session_token)
|
||||
if not owns_connector_session(session, user_id, source_type):
|
||||
lookup = dict(data)
|
||||
if not (lookup.get('connection_id') or lookup.get('session_token')):
|
||||
if not source_connection:
|
||||
return make_response(
|
||||
jsonify({"success": False, "error": "connection_id is required"}),
|
||||
400,
|
||||
)
|
||||
lookup['connection_id'] = source_connection
|
||||
session = service.resolve_request_connection(user_id, source_type, lookup)
|
||||
if session is None:
|
||||
return make_response(
|
||||
jsonify({"success": False, "error": "Invalid or unauthorized session"}),
|
||||
401,
|
||||
)
|
||||
else:
|
||||
# A grantee can't name a connection: theirs would read the
|
||||
# source's files with another account.
|
||||
with db_readonly() as conn:
|
||||
session = _owner_connector_session(conn, owner_id, source_type)
|
||||
session = _owner_connector_session(conn, owner_id, source_type, source_connection)
|
||||
if session is None:
|
||||
message = (
|
||||
"The owner needs to reconnect this source's account "
|
||||
@@ -629,7 +674,6 @@ class ConnectorSync(Resource):
|
||||
jsonify({"success": False, "error": message, "message": message}),
|
||||
409,
|
||||
)
|
||||
session_token = session["session_token"]
|
||||
|
||||
# Extract configuration from remote_data
|
||||
file_ids = remote_data.get('file_ids', [])
|
||||
@@ -641,7 +685,7 @@ class ConnectorSync(Resource):
|
||||
job_name=source.get('name'),
|
||||
user=owner_id,
|
||||
source_type=source_type,
|
||||
session_token=session_token,
|
||||
connection_id=str(session["id"]),
|
||||
file_ids=file_ids,
|
||||
folder_ids=folder_ids,
|
||||
recursive=recursive,
|
||||
|
||||
@@ -224,6 +224,7 @@ RULES: dict[tuple[str, str], Rule] = {
|
||||
("/api/get_chunks", "GET"): _rule("sources:read", (QUERY, "id")),
|
||||
("/api/sources/<string:source_id>/wiki/pages", "GET"): _rule("sources:read", (VIEW, "source_id")),
|
||||
("/api/sources/<string:source_id>/wiki/page", "GET"): _rule("sources:read", (VIEW, "source_id")),
|
||||
("/api/sources/<string:source_id>/wiki/settings", "GET"): _rule("sources:read", (VIEW, "source_id")),
|
||||
("/api/sources/<string:source_id>/graph", "GET"): _rule("sources:read", (VIEW, "source_id")),
|
||||
("/api/sources/<string:source_id>/graph/node/<string:node_id>", "GET"): _rule(
|
||||
"sources:read", (VIEW, "source_id")
|
||||
@@ -357,12 +358,15 @@ DENIED: dict[str, tuple[str, ...]] = {
|
||||
"/api/teams/<string:team_id>/grants": ("POST", "DELETE"),
|
||||
"/api/teams/<string:team_id>/transfer_owner": ("*",),
|
||||
"/api/resource_settings": ("PUT",),
|
||||
# Who may edit a wiki from outside the app is the owner's call in a session.
|
||||
"/api/sources/<string:source_id>/wiki/settings": ("PUT",),
|
||||
"/swagger.json": ("*",),
|
||||
}
|
||||
DENIED_PREFIXES = (
|
||||
"/api/admin/",
|
||||
"/api/auth/oidc/",
|
||||
"/api/connectors/",
|
||||
"/api/connections",
|
||||
"/api/devices",
|
||||
"/scim/",
|
||||
"/static/",
|
||||
|
||||
@@ -38,6 +38,7 @@ from docsgpt.agents.default_tools import (
|
||||
from docsgpt.api import api
|
||||
from docsgpt.api.pat.rules import allowed_ids
|
||||
from docsgpt.api.user.resource_access import AccessDenied, require
|
||||
from docsgpt.connectors.resolve import carry_removed_connection
|
||||
from docsgpt.core.model_utils import validate_model_id
|
||||
from docsgpt.core.url_validation import SSRFError, validate_url
|
||||
from docsgpt.security.safe_url import UnsafeUserUrlError, validate_user_base_url
|
||||
@@ -1138,7 +1139,8 @@ def _create_tool_from_spec(conn, user: str, tool: dict, secrets: dict, warnings:
|
||||
warnings.append(f"Tool type '{tool_type}' not available on this instance; skipped")
|
||||
return None
|
||||
config_requirements = inst.get_config_requirements() or {}
|
||||
config = dict(tool.get("config") or {})
|
||||
# An imported tool starts with no note that a connection was removed.
|
||||
config = carry_removed_connection(dict(tool.get("config") or {}), None)
|
||||
config.update(secrets or {})
|
||||
if tool_type == "api_tool":
|
||||
label = tool.get("display_name") or tool.get("name") or tool_type
|
||||
@@ -1491,18 +1493,24 @@ def _apply_workflow(
|
||||
if existing_id:
|
||||
existing = wf_repo.get(str(existing_id), user)
|
||||
if existing is not None:
|
||||
from docsgpt.api.user.resource_access import prune_sponsors
|
||||
from docsgpt.api.user.workflows.routes import _node_refs
|
||||
|
||||
pg_workflow_id = str(existing["id"])
|
||||
next_version = get_workflow_graph_version(existing) + 1
|
||||
_write_graph(conn, pg_workflow_id, next_version, nodes_data, edges_data)
|
||||
wf_repo.update(
|
||||
pg_workflow_id,
|
||||
user,
|
||||
{
|
||||
"name": name,
|
||||
"description": description,
|
||||
"current_graph_version": next_version,
|
||||
},
|
||||
)
|
||||
workflow_fields = {
|
||||
"name": name,
|
||||
"description": description,
|
||||
"current_graph_version": next_version,
|
||||
}
|
||||
# Sponsors of node resources the file dropped go too, so a stale
|
||||
# record can't vouch for the resource if someone adds it back.
|
||||
sponsors = existing.get("resource_sponsors") or {}
|
||||
pruned = prune_sponsors(sponsors, _node_refs(nodes_data))
|
||||
if pruned != sponsors:
|
||||
workflow_fields["resource_sponsors"] = pruned
|
||||
wf_repo.update(pg_workflow_id, user, workflow_fields)
|
||||
WorkflowNodesRepository(conn).delete_other_versions(pg_workflow_id, next_version)
|
||||
WorkflowEdgesRepository(conn).delete_other_versions(pg_workflow_id, next_version)
|
||||
return pg_workflow_id
|
||||
@@ -1512,6 +1520,25 @@ def _apply_workflow(
|
||||
return str(created["id"])
|
||||
|
||||
|
||||
def _prune_agent_sponsors(agents_repo: AgentsRepository, agent_id: str, user: str) -> None:
|
||||
"""Drop sponsor records of resources the imported agent no longer references.
|
||||
|
||||
Args:
|
||||
agents_repo: Repository on the import's connection.
|
||||
agent_id: The updated agent.
|
||||
user: Its owner.
|
||||
"""
|
||||
from docsgpt.api.user.resource_access import agent_refs, prune_sponsors
|
||||
|
||||
row = agents_repo.get(agent_id, user)
|
||||
if not row:
|
||||
return
|
||||
sponsors = row.get("resource_sponsors") or {}
|
||||
pruned = prune_sponsors(sponsors, agent_refs(row))
|
||||
if pruned != sponsors:
|
||||
agents_repo.update(agent_id, user, {"resource_sponsors": pruned})
|
||||
|
||||
|
||||
def apply_import(conn, user: str, doc: dict, resolution: Optional[dict] = None) -> dict:
|
||||
"""Create or update an agent from a parsed YAML doc.
|
||||
|
||||
@@ -1648,6 +1675,7 @@ def apply_import(conn, user: str, doc: dict, resolution: Optional[dict] = None)
|
||||
# Only a brand-new agent (the create path below) starts as a draft.
|
||||
fields = {**authoritative, **optional, "name": spec.get("name")}
|
||||
if agents_repo.update(str(target["agent_id"]), user, fields):
|
||||
_prune_agent_sponsors(agents_repo, str(target["agent_id"]), user)
|
||||
if orphaned_workflow_id and not agents_repo.count_by_workflow(
|
||||
orphaned_workflow_id, user
|
||||
):
|
||||
|
||||
@@ -29,13 +29,19 @@ from docsgpt.storage.db.base_repository import canonical_uuid, looks_like_uuid
|
||||
from docsgpt.api.user.resource_access import (
|
||||
AccessDenied,
|
||||
agent_refs,
|
||||
best_effort,
|
||||
cached_resolves,
|
||||
delete_settings,
|
||||
parse_confirmations,
|
||||
payload_for,
|
||||
plan_sponsors,
|
||||
require,
|
||||
resolve,
|
||||
resource_states,
|
||||
settings_many,
|
||||
sponsor_audience,
|
||||
sponsor_details,
|
||||
sponsors_after_save,
|
||||
sponsor_refusal,
|
||||
)
|
||||
from docsgpt.api.user.team_sharing import (
|
||||
can_access,
|
||||
@@ -483,6 +489,27 @@ def _build_create_kwargs(data: dict, *, image_url: str, agent_type: str) -> dict
|
||||
return kwargs
|
||||
|
||||
|
||||
def keep_owner_only_config(config: dict, existing_agent: dict, is_team_editor: bool) -> dict:
|
||||
"""Keep the parts of an agent's config only its owner may change.
|
||||
|
||||
``api_write_allowlist`` lets anyone with the agent's API key act on the
|
||||
owner's connected accounts, so a team editor's update keeps the stored
|
||||
value whatever it sends.
|
||||
|
||||
Args:
|
||||
config: The normalized config from the request.
|
||||
existing_agent: The agent row being updated.
|
||||
is_team_editor: The caller edits through a team share, not as owner.
|
||||
|
||||
Returns:
|
||||
The config to store.
|
||||
"""
|
||||
if not is_team_editor:
|
||||
return config
|
||||
stored = AgentConfig.parse(existing_agent.get("config")).api_write_allowlist
|
||||
return {**config, "api_write_allowlist": stored}
|
||||
|
||||
|
||||
def normalize_agent_config(raw):
|
||||
"""Validate an inbound ``config`` payload, returning the normalized dict.
|
||||
|
||||
@@ -537,15 +564,29 @@ class GetAgent(Resource):
|
||||
user = decoded_token["sub"]
|
||||
agent = None
|
||||
sponsored: list = []
|
||||
with db_readonly() as conn:
|
||||
states: list = []
|
||||
audience = None
|
||||
with db_readonly() as conn, cached_resolves():
|
||||
# Anyone who can see the agent reads it (a viewer needs it to
|
||||
# chat); what they get back is trimmed by their actions.
|
||||
ra = resolve(conn, "agent", agent_id, user)
|
||||
if ra is not None:
|
||||
agent = AgentsRepository(conn).get_by_id(ra.resource_id)
|
||||
# Edit-page detail: who vouches for resources the owner can't use.
|
||||
if agent and ra.can("view"):
|
||||
sponsored = sponsor_details(conn, "agent", agent)
|
||||
# Edit-page detail: who vouches for resources the owner can't
|
||||
# use, and which attached resources stopped running and why,
|
||||
# with their names. Only for people who may edit the agent:
|
||||
# the names can be an editor's private resources. Run state
|
||||
# never fails the read.
|
||||
if agent and ra.can("edit"):
|
||||
sponsored = sponsor_details(conn, "agent", agent, viewer=user)
|
||||
states = best_effort(
|
||||
conn, "the agent's resource states",
|
||||
lambda: resource_states(conn, "agent", agent, agent_refs(agent), user), [],
|
||||
)
|
||||
audience = best_effort(
|
||||
conn, "the agent's audience",
|
||||
lambda: sponsor_audience(conn, "agent", agent, states, sponsored), None,
|
||||
)
|
||||
if not agent:
|
||||
return {"status": "Not found"}, 404
|
||||
is_owner = ra.access == "owner"
|
||||
@@ -557,6 +598,9 @@ class GetAgent(Resource):
|
||||
access=ra.payload(),
|
||||
)
|
||||
data["resource_sponsors"] = sponsored
|
||||
data["resource_states"] = states
|
||||
if audience is not None:
|
||||
data["sponsor_audience"] = audience
|
||||
return make_response(jsonify(data), 200)
|
||||
except Exception as e:
|
||||
current_app.logger.error(f"Agent fetch error: {e}", exc_info=True)
|
||||
@@ -1035,14 +1079,11 @@ class UpdateAgent(Resource):
|
||||
)
|
||||
pg_agent_id = str(existing_agent["id"])
|
||||
existing_image = existing_agent.get("image", "") or ""
|
||||
image_url, image_error = handle_image_upload(
|
||||
request,
|
||||
existing_image,
|
||||
existing_agent.get("user_id") or user,
|
||||
storage,
|
||||
)
|
||||
if image_error:
|
||||
return image_error
|
||||
# The image is stored only once the save is known to go
|
||||
# ahead, so a refused save (a sponsor confirmation round
|
||||
# trip, a validation error) leaves no orphaned file.
|
||||
image_file = request.files.get("image")
|
||||
has_new_image = bool(image_file and image_file.filename)
|
||||
|
||||
update_fields: dict = {}
|
||||
allowed_fields = [
|
||||
@@ -1153,7 +1194,9 @@ class UpdateAgent(Resource):
|
||||
exc,
|
||||
)
|
||||
return _reject(INVALID_CONFIG_MESSAGE, user, field)
|
||||
update_fields["config"] = normalized_config or {}
|
||||
update_fields["config"] = keep_owner_only_config(
|
||||
normalized_config or {}, existing_agent, is_team_editor,
|
||||
)
|
||||
elif field == "limited_token_mode":
|
||||
raw_value = data.get("limited_token_mode", False)
|
||||
bool_value = (
|
||||
@@ -1288,9 +1331,7 @@ class UpdateAgent(Resource):
|
||||
f"Field '{field}' cannot be empty", user, field
|
||||
)
|
||||
update_fields[field] = value
|
||||
if image_url and image_url != existing_image:
|
||||
update_fields["image"] = image_url
|
||||
if not update_fields:
|
||||
if not update_fields and not has_new_image:
|
||||
return _reject("No valid update data provided", user)
|
||||
|
||||
newly_generated_key = None
|
||||
@@ -1406,18 +1447,6 @@ class UpdateAgent(Resource):
|
||||
403,
|
||||
)
|
||||
|
||||
# A resource the owner can't use runs as the editor who
|
||||
# attached it (its sponsor); record who that is.
|
||||
after_save = dict(existing_agent)
|
||||
for ref_field in ("source_id", "extra_source_ids", "prompt_id", "tools"):
|
||||
if ref_field in update_fields:
|
||||
after_save[ref_field] = update_fields[ref_field]
|
||||
sponsors = sponsors_after_save(
|
||||
conn, "agent", existing_agent, owner_id, user, agent_refs(after_save)
|
||||
)
|
||||
if sponsors != (existing_agent.get("resource_sponsors") or {}):
|
||||
update_fields["resource_sponsors"] = sponsors
|
||||
|
||||
# Guardrails and the pooled quota are policy: an unchanged
|
||||
# value re-sent by a full-form save is fine, a change needs
|
||||
# ``edit_policy``.
|
||||
@@ -1426,6 +1455,46 @@ class UpdateAgent(Resource):
|
||||
AccessDenied(403, "Your access doesn't allow changing guardrails or limits")
|
||||
)
|
||||
|
||||
# A resource the owner can't use runs as the editor who
|
||||
# attached it (its sponsor). Only someone who owns or edits
|
||||
# it may sponsor it, and only after confirming: the save is
|
||||
# refused (409) until ``confirm_sponsor`` lists every
|
||||
# resource it would newly sponsor.
|
||||
after_save = dict(existing_agent)
|
||||
for ref_field in ("source_id", "extra_source_ids", "prompt_id", "tools"):
|
||||
if ref_field in update_fields:
|
||||
after_save[ref_field] = update_fields[ref_field]
|
||||
plan = plan_sponsors(
|
||||
conn,
|
||||
"agent",
|
||||
existing_agent,
|
||||
owner_id,
|
||||
user,
|
||||
agent_refs(after_save),
|
||||
previous_refs=agent_refs(existing_agent),
|
||||
confirmed=parse_confirmations(data.get("confirm_sponsor")),
|
||||
)
|
||||
refusal = sponsor_refusal(
|
||||
conn, "agent", existing_agent, plan, api_key=bool(update_fields.get("key")) or None
|
||||
)
|
||||
if refusal is not None:
|
||||
body, status = refusal
|
||||
return make_response(jsonify(body), status)
|
||||
if plan.sponsors != (existing_agent.get("resource_sponsors") or {}):
|
||||
update_fields["resource_sponsors"] = plan.sponsors
|
||||
|
||||
if has_new_image:
|
||||
image_url, image_error = handle_image_upload(
|
||||
request,
|
||||
existing_image,
|
||||
existing_agent.get("user_id") or user,
|
||||
storage,
|
||||
)
|
||||
if image_error:
|
||||
return image_error
|
||||
if image_url and image_url != existing_image:
|
||||
update_fields["image"] = image_url
|
||||
|
||||
# Apply update. Owner writes use the dual-key guard; team-editor
|
||||
# writes go by-id (already authorized) with an optimistic-lock
|
||||
# check when the client supplies the row's expected updated_at.
|
||||
|
||||
File diff suppressed because it is too large.
Load diff
@@ -233,6 +233,33 @@ def _scheduler_still_allowed(schedule: Dict[str, Any], agent_config: Dict[str, A
|
||||
return resolve(conn, "agent", str(agent_id), user_id) is not None
|
||||
|
||||
|
||||
def _schedule_caller_rules(schedule: Dict[str, Any], agent_config: Dict[str, Any]) -> Dict[str, bool]:
|
||||
"""The caller rules a scheduled run keeps, though it acts as the owner.
|
||||
|
||||
A schedule set through the agent's API key is an external caller's; one
|
||||
a user without a team grant set (they reach the agent by its public
|
||||
link) is a public-link caller's. Neither can approve for the owner, so
|
||||
the run writes on the owner's accounts and credentials only as far as
|
||||
the agent's API write allowlist allows.
|
||||
|
||||
Args:
|
||||
schedule: The schedule row.
|
||||
agent_config: The agent row (or the agentless ephemeral config).
|
||||
|
||||
Returns:
|
||||
``external_caller`` and ``public_link_caller`` flags for the run.
|
||||
"""
|
||||
public = False
|
||||
user_id = schedule.get("user_id")
|
||||
agent_id = agent_config.get("id")
|
||||
if agent_id and user_id and agent_config.get("user_id") != user_id:
|
||||
from docsgpt.api.user.resource_access import resolve
|
||||
|
||||
with get_engine().connect() as conn:
|
||||
public = resolve(conn, "agent", str(agent_id), user_id) is None
|
||||
return {"external_caller": schedule.get("created_via") == "api", "public_link_caller": public}
|
||||
|
||||
|
||||
def execute_scheduled_run_body(run_id: str, celery_task_id: Optional[str]) -> Dict[str, Any]:
|
||||
"""Execute one scheduled run by id; returns a result dict for tracing."""
|
||||
if not settings.POSTGRES_URI:
|
||||
@@ -324,6 +351,7 @@ def execute_scheduled_run_body(run_id: str, celery_task_id: Optional[str]) -> Di
|
||||
# to the user who scheduled it (not always the agent's owner).
|
||||
request_id=str(run_id),
|
||||
trace_user_id=run.get("user_id") or schedule.get("user_id"),
|
||||
**_schedule_caller_rules(schedule, agent_config),
|
||||
)
|
||||
except SoftTimeLimitExceeded:
|
||||
timed_out = True
|
||||
|
||||
@@ -74,6 +74,12 @@ def _get_provider_from_remote_data(remote_data):
|
||||
return None
|
||||
|
||||
|
||||
def _connection_id(row: dict) -> str | None:
|
||||
"""The connection a source syncs from, as a string id, or None."""
|
||||
value = row.get("connection_id")
|
||||
return str(value) if value else None
|
||||
|
||||
|
||||
def _with_access(entry: dict, access: Optional[str], switches: Optional[dict]) -> dict:
|
||||
"""Add ``access`` + ``allowed_actions`` to a listed source row.
|
||||
|
||||
@@ -150,6 +156,7 @@ class CombinedJson(Resource):
|
||||
"config": SourceConfig.parse(index.get("config")).model_dump(),
|
||||
"ownership": ownership,
|
||||
"team_access": team_access,
|
||||
"connectionId": _connection_id(index),
|
||||
}
|
||||
return _with_access(
|
||||
entry,
|
||||
@@ -247,6 +254,7 @@ class PaginatedSources(Resource):
|
||||
"team_access": (
|
||||
None if owned else team_shared.get(str(doc["id"]))
|
||||
),
|
||||
"connectionId": _connection_id(doc),
|
||||
}
|
||||
paginated_docs.append(
|
||||
_with_access(
|
||||
@@ -269,6 +277,70 @@ class PaginatedSources(Resource):
|
||||
return make_response(jsonify({"success": False}), 400)
|
||||
|
||||
|
||||
def delete_source(owner: str, doc: dict, *, actor: Optional[str] = None) -> bool:
|
||||
"""Delete a source's index, stored files and row. Returns whether it worked.
|
||||
|
||||
Args:
|
||||
owner: The source's owner; the row is deleted as them.
|
||||
doc: The source row, already authorised for the caller.
|
||||
actor: Who asked for the delete (a team editor, say), when not the
|
||||
owner. Recorded on the audit event.
|
||||
"""
|
||||
actor = actor or owner
|
||||
storage = StorageCreator.get_storage()
|
||||
resolved_id = str(doc["id"])
|
||||
source_id = resolved_id
|
||||
|
||||
try:
|
||||
if settings.VECTOR_STORE == "faiss":
|
||||
index_path = f"indexes/{resolved_id}"
|
||||
# index.pkl is the legacy sidecar; index.json the current one.
|
||||
# Older sources have only the former, so clear whichever exist.
|
||||
for index_file in ("index.faiss", "index.json", "index.pkl"):
|
||||
if storage.file_exists(f"{index_path}/{index_file}"):
|
||||
storage.delete_file(f"{index_path}/{index_file}")
|
||||
else:
|
||||
vectorstore = VectorCreator.create_vectorstore(
|
||||
settings.VECTOR_STORE, source_id=source_id
|
||||
)
|
||||
vectorstore.delete_index()
|
||||
if "file_path" in doc and doc["file_path"]:
|
||||
file_path = doc["file_path"]
|
||||
if storage.is_directory(file_path):
|
||||
files = storage.list_files(file_path)
|
||||
for f in files:
|
||||
storage.delete_file(f)
|
||||
else:
|
||||
storage.delete_file(file_path)
|
||||
except FileNotFoundError:
|
||||
pass
|
||||
except Exception as err:
|
||||
current_app.logger.error(
|
||||
f"Error deleting files and indexes: {err}", exc_info=True
|
||||
)
|
||||
return False
|
||||
try:
|
||||
with db_session() as conn:
|
||||
# The AFTER DELETE trigger drops the source's team grants; the
|
||||
# owner switches have no FK, so clear them here.
|
||||
SourcesRepository(conn).delete(resolved_id, owner)
|
||||
delete_settings(conn, "source", resolved_id)
|
||||
record_event(
|
||||
conn,
|
||||
"source.deleted",
|
||||
actor=actor,
|
||||
source_id=resolved_id,
|
||||
name=doc.get("name"),
|
||||
owner=owner if owner != actor else None,
|
||||
)
|
||||
except Exception as err:
|
||||
current_app.logger.error(
|
||||
f"Error deleting source row: {err}", exc_info=True
|
||||
)
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
@sources_ns.route("/delete_old")
|
||||
class DeleteOldIndexes(Resource):
|
||||
@api.doc(
|
||||
@@ -295,56 +367,7 @@ class DeleteOldIndexes(Resource):
|
||||
except Exception as err:
|
||||
current_app.logger.error(f"Error looking up source: {err}", exc_info=True)
|
||||
return make_response(jsonify({"success": False}), 400)
|
||||
owner = ra.owner_id
|
||||
storage = StorageCreator.get_storage()
|
||||
resolved_id = str(doc["id"])
|
||||
|
||||
try:
|
||||
if settings.VECTOR_STORE == "faiss":
|
||||
index_path = f"indexes/{resolved_id}"
|
||||
# index.pkl is the legacy sidecar; index.json the current one.
|
||||
# Older sources have only the former, so clear whichever exist.
|
||||
for index_file in ("index.faiss", "index.json", "index.pkl"):
|
||||
if storage.file_exists(f"{index_path}/{index_file}"):
|
||||
storage.delete_file(f"{index_path}/{index_file}")
|
||||
else:
|
||||
vectorstore = VectorCreator.create_vectorstore(
|
||||
settings.VECTOR_STORE, source_id=resolved_id
|
||||
)
|
||||
vectorstore.delete_index()
|
||||
if "file_path" in doc and doc["file_path"]:
|
||||
file_path = doc["file_path"]
|
||||
if storage.is_directory(file_path):
|
||||
files = storage.list_files(file_path)
|
||||
for f in files:
|
||||
storage.delete_file(f)
|
||||
else:
|
||||
storage.delete_file(file_path)
|
||||
except FileNotFoundError:
|
||||
pass
|
||||
except Exception as err:
|
||||
current_app.logger.error(
|
||||
f"Error deleting files and indexes: {err}", exc_info=True
|
||||
)
|
||||
return make_response(jsonify({"success": False}), 400)
|
||||
try:
|
||||
with db_session() as conn:
|
||||
# The AFTER DELETE trigger drops the source's team grants; the
|
||||
# owner switches have no FK, so clear them here.
|
||||
SourcesRepository(conn).delete(resolved_id, owner)
|
||||
delete_settings(conn, "source", resolved_id)
|
||||
record_event(
|
||||
conn,
|
||||
"source.deleted",
|
||||
actor=user,
|
||||
source_id=resolved_id,
|
||||
name=doc.get("name"),
|
||||
owner=owner if owner != user else None,
|
||||
)
|
||||
except Exception as err:
|
||||
current_app.logger.error(
|
||||
f"Error deleting source row: {err}", exc_info=True
|
||||
)
|
||||
if not delete_source(ra.owner_id, doc, actor=user):
|
||||
return make_response(jsonify({"success": False}), 400)
|
||||
return make_response(jsonify({"success": True}), 200)
|
||||
|
||||
@@ -467,6 +490,8 @@ class SyncSource(Resource):
|
||||
sync_frequency=doc.get("sync_frequency", "never"),
|
||||
retriever=doc.get("retriever", "classic"),
|
||||
doc_id=str(doc["id"]),
|
||||
# S3 and GitHub sources made from a connection read with its keys.
|
||||
connection_id=str(doc["connection_id"]) if doc.get("connection_id") else None,
|
||||
)
|
||||
except Exception as err:
|
||||
current_app.logger.error(
|
||||
@@ -1031,6 +1056,90 @@ class WikiPage(Resource):
|
||||
)
|
||||
|
||||
|
||||
def _wiki_settings_body(doc: dict, ra) -> dict:
|
||||
"""The wiki settings response: the stored switch plus the caller's access."""
|
||||
return {
|
||||
"success": True,
|
||||
"allow_outside_edits": bool(doc.get("wiki_outside_edits")),
|
||||
**ra.payload(),
|
||||
}
|
||||
|
||||
|
||||
@sources_ns.route("/sources/<string:source_id>/wiki/settings")
|
||||
class WikiSettings(Resource):
|
||||
@api.doc(
|
||||
description="A wiki's settings. Anyone who can see the wiki may read "
|
||||
"them; returns allow_outside_edits plus the caller's access."
|
||||
)
|
||||
def get(self, source_id):
|
||||
decoded_token = request.decoded_token
|
||||
if not decoded_token:
|
||||
return make_response(jsonify({"success": False}), 401)
|
||||
user = decoded_token.get("sub")
|
||||
try:
|
||||
with db_readonly() as conn:
|
||||
try:
|
||||
doc, ra = load_source(conn, source_id, user, "use")
|
||||
except AccessDenied as err:
|
||||
return denied_response(err)
|
||||
except Exception as err:
|
||||
current_app.logger.error(
|
||||
f"Error reading wiki settings for {source_id}: {err}", exc_info=True
|
||||
)
|
||||
return make_response(jsonify({"success": False}), 400)
|
||||
return make_response(jsonify(_wiki_settings_body(doc, ra)), 200)
|
||||
|
||||
@api.doc(
|
||||
description="Change a wiki's settings (owner only, manage_settings). "
|
||||
"Body: {\"allow_outside_edits\": bool}: whether runs from an agent's "
|
||||
"API key or widget may edit the wiki."
|
||||
)
|
||||
def put(self, source_id):
|
||||
decoded_token = request.decoded_token
|
||||
if not decoded_token:
|
||||
return make_response(jsonify({"success": False}), 401)
|
||||
user = decoded_token.get("sub")
|
||||
data = request.get_json(silent=True) or {}
|
||||
allowed = data.get("allow_outside_edits")
|
||||
if not isinstance(allowed, bool):
|
||||
return make_response(
|
||||
jsonify(
|
||||
{"success": False, "message": "allow_outside_edits must be true or false"}
|
||||
),
|
||||
400,
|
||||
)
|
||||
try:
|
||||
with db_session() as conn:
|
||||
try:
|
||||
doc, ra = load_source(conn, source_id, user, "manage_settings")
|
||||
except AccessDenied as err:
|
||||
return denied_response(err)
|
||||
if SourceConfig.parse(doc.get("config")).kind != "wiki":
|
||||
return make_response(
|
||||
jsonify({"success": False, "message": "Source is not a wiki"}), 400
|
||||
)
|
||||
if not SourcesRepository(conn).set_wiki_outside_edits(
|
||||
str(doc["id"]), ra.owner_id, allowed
|
||||
):
|
||||
return make_response(
|
||||
jsonify({"success": False, "message": "Source not found"}), 404
|
||||
)
|
||||
record_event(
|
||||
conn,
|
||||
"source.wiki_settings_updated",
|
||||
actor=user,
|
||||
source_id=str(doc["id"]),
|
||||
allow_outside_edits=allowed,
|
||||
)
|
||||
doc["wiki_outside_edits"] = allowed
|
||||
except Exception as err:
|
||||
current_app.logger.error(
|
||||
f"Error updating wiki settings for {source_id}: {err}", exc_info=True
|
||||
)
|
||||
return make_response(jsonify({"success": False}), 400)
|
||||
return make_response(jsonify(_wiki_settings_body(doc, ra)), 200)
|
||||
|
||||
|
||||
def _source_is_blank(doc):
|
||||
"""True when a source has no ingested files to convert into pages."""
|
||||
structure = doc.get("directory_structure") or {}
|
||||
|
||||
@@ -29,7 +29,6 @@ from docsgpt.security.zip_archive import (
|
||||
)
|
||||
from docsgpt.storage.db.repositories.connector_sessions import (
|
||||
ConnectorSessionsRepository,
|
||||
owns_connector_session,
|
||||
)
|
||||
from docsgpt.storage.db.repositories.idempotency import IdempotencyRepository
|
||||
from docsgpt.storage.db.repositories.sources import SourcesRepository
|
||||
@@ -452,6 +451,45 @@ class UploadFile(Resource):
|
||||
return make_response(jsonify(response_payload), 200)
|
||||
|
||||
|
||||
def _remote_credentials(user, source, config):
|
||||
"""Split an S3 / Reddit request into loader config and the connection holding its keys.
|
||||
|
||||
A request naming a ``connection_id`` uses that connection's stored keys.
|
||||
A request carrying keys (the form before connections) stores them on a
|
||||
connection, so they are entered once and never land in
|
||||
``sources.remote_data``. When a multi-user install still runs on the
|
||||
public default encryption key, the keys stay with the source as before.
|
||||
|
||||
Returns:
|
||||
``(source_data, connection_id, error_response)``.
|
||||
"""
|
||||
from docsgpt.connectors import catalog, service
|
||||
|
||||
definition = catalog.get_definition(source)
|
||||
credential_keys = {f.key for f in definition.credential_fields}
|
||||
public = {k: v for k, v in config.items() if k not in credential_keys and k != "connection_id"}
|
||||
connection_id = config.get("connection_id")
|
||||
if connection_id:
|
||||
with db_readonly() as conn:
|
||||
row = ConnectorSessionsRepository(conn).get_for_user(str(connection_id), user)
|
||||
if row is None or catalog.connector_key_for_row(row) != source:
|
||||
return None, None, make_response(
|
||||
jsonify({"success": False, "error": "Invalid or unauthorized connection"}), 401,
|
||||
)
|
||||
return public, str(row["id"]), None
|
||||
provided = {k: config[k] for k in credential_keys if config.get(k) not in (None, "")}
|
||||
if not provided:
|
||||
return config, None, None
|
||||
try:
|
||||
with db_session() as conn:
|
||||
row, _ = service.create_api_key_connection(conn, user, definition, provided)
|
||||
except service.ConnectorDisabled as err:
|
||||
return None, None, make_response(jsonify({"success": False, "error": str(err)}), 403)
|
||||
except (service.EncryptionKeyNotConfigured, ValueError):
|
||||
return config, None, None
|
||||
return public, str(row["id"]), None
|
||||
|
||||
|
||||
@sources_upload_ns.route("/remote")
|
||||
class UploadRemote(Resource):
|
||||
@api.expect(
|
||||
@@ -521,32 +559,35 @@ class UploadRemote(Resource):
|
||||
try:
|
||||
config = json.loads(data["data"])
|
||||
source_data = None
|
||||
connection_id = None
|
||||
|
||||
if data["source"] == "github":
|
||||
source_data = config.get("repo_url")
|
||||
elif data["source"] in ["crawler", "url", "sitemap"]:
|
||||
source_data = config.get("url")
|
||||
elif data["source"] == "reddit":
|
||||
source_data = config
|
||||
elif data["source"] == "s3":
|
||||
source_data = config
|
||||
elif data["source"] in ("reddit", "s3"):
|
||||
source_data, connection_id, error = _remote_credentials(user, data["source"], config)
|
||||
if error is not None:
|
||||
if scoped_key:
|
||||
_release_claim(scoped_key)
|
||||
return error
|
||||
elif data["source"] in ConnectorCreator.get_supported_connectors():
|
||||
session_token = config.get("session_token")
|
||||
if not session_token:
|
||||
if not (config.get("connection_id") or config.get("session_token")):
|
||||
if scoped_key:
|
||||
_release_claim(scoped_key)
|
||||
return make_response(
|
||||
jsonify(
|
||||
{
|
||||
"success": False,
|
||||
"error": f"Missing session_token in {data['source']} configuration",
|
||||
"error": f"Missing connection_id in {data['source']} configuration",
|
||||
}
|
||||
),
|
||||
400,
|
||||
)
|
||||
with db_readonly() as conn:
|
||||
connector_session = ConnectorSessionsRepository(conn).get_by_session_token(session_token)
|
||||
if not owns_connector_session(connector_session, user, data["source"]):
|
||||
from docsgpt.connectors import service as connection_service
|
||||
|
||||
connector_session = connection_service.resolve_request_connection(user, data["source"], config)
|
||||
if connector_session is None:
|
||||
if scoped_key:
|
||||
_release_claim(scoped_key)
|
||||
return make_response(
|
||||
@@ -577,7 +618,7 @@ class UploadRemote(Resource):
|
||||
"job_name": data["name"],
|
||||
"user": user,
|
||||
"source_type": data["source"],
|
||||
"session_token": session_token,
|
||||
"connection_id": str(connector_session["id"]),
|
||||
"file_ids": file_ids,
|
||||
"folder_ids": folder_ids,
|
||||
"recursive": config.get("recursive", False),
|
||||
@@ -612,6 +653,7 @@ class UploadRemote(Resource):
|
||||
remote_kwargs = {
|
||||
"kwargs": {
|
||||
"source_data": source_data,
|
||||
"connection_id": connection_id,
|
||||
"job_name": data["name"],
|
||||
"user": user,
|
||||
"loader": data["source"],
|
||||
|
||||
@@ -79,6 +79,12 @@ DURABLE_TASK = dict(
|
||||
)
|
||||
|
||||
|
||||
def _connection_unavailable():
|
||||
from docsgpt.connectors.service import ConnectionUnavailable
|
||||
|
||||
return ConnectionUnavailable
|
||||
|
||||
|
||||
def durable_task(**overrides) -> Dict:
|
||||
"""Return ``DURABLE_TASK`` with per-task overrides applied.
|
||||
|
||||
@@ -167,13 +173,16 @@ def ingest(
|
||||
@with_idempotency(task_name="ingest_remote", on_poison=_emit_ingest_poison_event)
|
||||
def ingest_remote(
|
||||
self, source_data, job_name, user, loader,
|
||||
config=None, idempotency_key=None, source_id=None,
|
||||
config=None, idempotency_key=None, source_id=None, connection_id=None,
|
||||
sync_frequency="never",
|
||||
):
|
||||
resp = remote_worker(
|
||||
self, source_data, job_name, user, loader,
|
||||
sync_frequency=sync_frequency,
|
||||
config=config,
|
||||
idempotency_key=idempotency_key,
|
||||
source_id=source_id,
|
||||
connection_id=connection_id,
|
||||
)
|
||||
return resp
|
||||
|
||||
@@ -249,6 +258,15 @@ def schedule_syncs(self, frequency):
|
||||
return resp
|
||||
|
||||
|
||||
@celery.task(bind=True, acks_late=True, autoretry_for=(Exception,), max_retries=3, retry_backoff=60,
|
||||
dont_autoretry_for=(_connection_unavailable(),))
|
||||
def sync_connector_source(self, source_id):
|
||||
"""Re-sync one connector source from its connection, with no browser involved."""
|
||||
from docsgpt.worker import sync_connector_source as run
|
||||
|
||||
return run(self, source_id)
|
||||
|
||||
|
||||
@celery.task(bind=True)
|
||||
def sync_source(
|
||||
self,
|
||||
@@ -259,6 +277,7 @@ def sync_source(
|
||||
sync_frequency,
|
||||
retriever,
|
||||
doc_id,
|
||||
connection_id=None,
|
||||
):
|
||||
resp = sync(
|
||||
self,
|
||||
@@ -269,6 +288,7 @@ def sync_source(
|
||||
sync_frequency,
|
||||
retriever,
|
||||
doc_id,
|
||||
connection_id=connection_id,
|
||||
)
|
||||
return resp
|
||||
|
||||
@@ -408,7 +428,9 @@ except Exception:
|
||||
pass
|
||||
|
||||
|
||||
@celery.task(**DURABLE_TASK)
|
||||
# A revoked or disconnected connection will not heal by retrying; the
|
||||
# service has already paused the source and told its owner to reconnect.
|
||||
@celery.task(**durable_task(dont_autoretry_for=(DocumentParseError, _connection_unavailable())))
|
||||
@with_idempotency(
|
||||
task_name="ingest_connector_task", on_poison=_emit_ingest_poison_event,
|
||||
)
|
||||
@@ -428,6 +450,7 @@ def ingest_connector_task(
|
||||
config=None,
|
||||
idempotency_key=None,
|
||||
source_id=None,
|
||||
connection_id=None,
|
||||
):
|
||||
from docsgpt.worker import ingest_connector
|
||||
|
||||
@@ -437,6 +460,7 @@ def ingest_connector_task(
|
||||
user,
|
||||
source_type,
|
||||
session_token=session_token,
|
||||
connection_id=connection_id,
|
||||
file_ids=file_ids,
|
||||
folder_ids=folder_ids,
|
||||
recursive=recursive,
|
||||
|
||||
+210
-49
@@ -5,6 +5,7 @@ from urllib.parse import urlencode, urlparse
|
||||
from flask import current_app, jsonify, make_response, redirect, request
|
||||
from flask_restx import Namespace, Resource, fields
|
||||
|
||||
from docsgpt.agents.tool_pins import carry_pins_between
|
||||
from docsgpt.agents.tools.mcp_tool import MCPOAuthManager, MCPTool
|
||||
from docsgpt.api import api
|
||||
from docsgpt.api.user.resource_access import AccessDenied, require
|
||||
@@ -18,6 +19,7 @@ from docsgpt.api.user.tools.routes import (
|
||||
transform_actions,
|
||||
)
|
||||
from docsgpt.cache import get_redis_instance
|
||||
from docsgpt.connectors.resolve import REMOVED_CONNECTION_KEY
|
||||
from docsgpt.core.url_validation import SSRFError, validate_url
|
||||
from docsgpt.security.encryption import decrypt_credentials, encrypt_credentials
|
||||
from docsgpt.storage.db.repositories.connector_sessions import (
|
||||
@@ -36,13 +38,19 @@ def _sanitize_mcp_transport(config):
|
||||
"""Normalise and validate the transport_type field.
|
||||
|
||||
Strips ``command`` / ``args`` keys that are only valid for local STDIO
|
||||
transports and returns the cleaned transport type string.
|
||||
transports, ``connection_id``, which only the tool executor sets
|
||||
(it picks whose MCP tokens the tool uses), and the note that a
|
||||
connection was removed, which only removing one writes (a save stores
|
||||
a fresh config, so it drops that note too). Returns the cleaned
|
||||
transport type string.
|
||||
"""
|
||||
transport_type = (config.get("transport_type") or "auto").lower()
|
||||
if transport_type not in _ALLOWED_TRANSPORTS:
|
||||
raise ValueError(f"Unsupported transport_type: {transport_type}")
|
||||
config.pop("command", None)
|
||||
config.pop("args", None)
|
||||
config.pop("connection_id", None)
|
||||
config.pop(REMOVED_CONNECTION_KEY, None)
|
||||
config["transport_type"] = transport_type
|
||||
return transport_type
|
||||
|
||||
@@ -84,14 +92,140 @@ def _validate_mcp_server_url(config: dict) -> None:
|
||||
raise ValueError(f"Invalid server URL: {exc}") from exc
|
||||
|
||||
|
||||
def _mcp_connection(user, config, auth_type, auth_credentials, display_name):
|
||||
"""The connection an MCP tool runs with; created on first save.
|
||||
|
||||
OAuth servers already have one (the sign-in stored its tokens there).
|
||||
Key, bearer and basic auth store their secret on a connection; servers
|
||||
with no auth get a credential-less connection so they still appear on
|
||||
the Connectors page. Returns None when a multi-user install runs on the
|
||||
default encryption key, which keeps the legacy per-tool secret.
|
||||
"""
|
||||
from docsgpt.connectors import catalog, service
|
||||
|
||||
base = catalog.base_url(config.get("server_url"))
|
||||
if not base:
|
||||
return None
|
||||
if auth_type == "oauth":
|
||||
with db_readonly() as conn:
|
||||
row = service._mcp_row(conn, user, base, None)
|
||||
return str(row["id"]) if row else None
|
||||
definition = catalog.get_definition("custom_mcp")
|
||||
host = base.split("://")[-1]
|
||||
try:
|
||||
with db_session() as conn:
|
||||
if auth_credentials:
|
||||
row, _ = service.create_api_key_connection(
|
||||
conn, user, definition, auth_credentials, server_url=base, display_name=display_name,
|
||||
)
|
||||
else:
|
||||
repo = ConnectorSessionsRepository(conn)
|
||||
row = repo.find_account(user, "custom_mcp", server_url=base, account_label=host) or repo.create(
|
||||
user, "custom_mcp", connector_key="custom_mcp", auth_kind="none",
|
||||
display_name=display_name, account_label=host, server_url=base,
|
||||
)
|
||||
except service.EncryptionKeyNotConfigured:
|
||||
return None
|
||||
return str(row["id"]) if row else None
|
||||
|
||||
|
||||
def _previous_connection(existing_doc, config, owner_id) -> str | None:
|
||||
"""The saved tool's connection, kept only while the server is unchanged.
|
||||
|
||||
An edit that cannot resolve a connection of its own (the default
|
||||
encryption key blocks storing a new secret) must not carry the old
|
||||
one over to a different server (its key would be sent there), to a
|
||||
connection the tool's owner does not own, or to one that signs in
|
||||
another way.
|
||||
"""
|
||||
from docsgpt.connectors import catalog
|
||||
|
||||
connection_id = (existing_doc or {}).get("connection_id")
|
||||
base = catalog.base_url(config.get("server_url"))
|
||||
if not connection_id or not base:
|
||||
return None
|
||||
with db_readonly() as conn:
|
||||
row = ConnectorSessionsRepository(conn).get_for_user(str(connection_id), owner_id)
|
||||
if row is None or catalog.base_url(row.get("server_url")) != base:
|
||||
return None
|
||||
wanted = {"oauth": "mcp_oauth", "none": "none"}.get(config.get("auth_type") or "none", "api_key")
|
||||
if (row.get("auth_kind") or "") != wanted:
|
||||
return None
|
||||
return str(connection_id)
|
||||
|
||||
|
||||
def _mcp_policy_error(config: dict):
|
||||
"""A 403 when an admin turned this MCP server's connector off, else None.
|
||||
|
||||
A preset's own switch applies to its server; any other server is a
|
||||
custom connector and needs "Allow custom MCP servers".
|
||||
"""
|
||||
from docsgpt.connectors import catalog, service
|
||||
|
||||
preset = catalog.preset_for_url(config.get("server_url"))
|
||||
key = preset.key if preset else "custom_mcp"
|
||||
try:
|
||||
with db_readonly() as conn:
|
||||
service.ensure_connector_allowed(conn, key)
|
||||
except service.ConnectorDisabled:
|
||||
return make_response(
|
||||
jsonify({"success": False, "error": "This MCP server is turned off by an admin", "code": "disabled"}),
|
||||
403,
|
||||
)
|
||||
except Exception:
|
||||
# Fail closed: a server whose admin switch cannot be read is not contacted.
|
||||
current_app.logger.warning("Could not read connector policies", exc_info=True)
|
||||
return make_response(
|
||||
jsonify({"success": False, "error": "Could not check whether this MCP server is allowed"}),
|
||||
503,
|
||||
)
|
||||
return None
|
||||
|
||||
|
||||
def _stored_mcp_credentials(existing_doc: dict, owner_id: str) -> dict:
|
||||
"""The secrets a saved MCP tool authenticates with, for reuse on an edit.
|
||||
|
||||
A legacy tool keeps them encrypted in its config; a connection-backed one
|
||||
on its key-based connection, read only while that connection is the
|
||||
owner's and was stored for the tool's own server.
|
||||
|
||||
Args:
|
||||
existing_doc: The stored ``user_tools`` row.
|
||||
owner_id: The tool's owner.
|
||||
|
||||
Returns:
|
||||
The stored credentials, or ``{}`` when there are none to reuse.
|
||||
"""
|
||||
from docsgpt.connectors import catalog, service
|
||||
|
||||
existing_config = existing_doc.get("config") or {}
|
||||
if existing_config.get("encrypted_credentials"):
|
||||
return decrypt_credentials(existing_config["encrypted_credentials"], owner_id)
|
||||
connection_id = existing_doc.get("connection_id")
|
||||
if not connection_id:
|
||||
return {}
|
||||
with db_readonly() as conn:
|
||||
row = ConnectorSessionsRepository(conn).get_for_user(str(connection_id), owner_id)
|
||||
if (
|
||||
row is None
|
||||
or (row.get("auth_kind") or "") != "api_key"
|
||||
or catalog.base_url(row.get("server_url")) != catalog.base_url(existing_config.get("server_url"))
|
||||
):
|
||||
return {}
|
||||
try:
|
||||
return service.get_credentials(row)
|
||||
except service.ConnectionUnavailable:
|
||||
return {}
|
||||
|
||||
|
||||
def _existing_mcp_context(tool_id, user, config):
|
||||
"""Resolve the stored MCP tool a test/save refers to, and its credentials.
|
||||
|
||||
With no ``tool_id`` the caller acts on their own new server. With one,
|
||||
the caller needs ``edit_credentials`` on that tool and everything runs as
|
||||
its owner. Stored secrets are write-only, so an empty secret field reuses
|
||||
the stored one while the origin (scheme, host, port) is unchanged; a new
|
||||
origin never inherits them.
|
||||
the stored one (on the tool, or on its key-based connection) while the
|
||||
origin (scheme, host, port) is unchanged; a new origin never inherits them.
|
||||
A server that is or would become OAuth is the owner's alone (its tokens
|
||||
are the owner's sign-in).
|
||||
|
||||
@@ -123,9 +257,9 @@ def _existing_mcp_context(tool_id, user, config):
|
||||
jsonify({"success": False, "message": _CREDENTIALS_FOR_NEW_SERVER}), 400
|
||||
)
|
||||
credentials = dict(auth_credentials)
|
||||
existing_encrypted = None if moved else existing_config.get("encrypted_credentials")
|
||||
if existing_encrypted:
|
||||
credentials = {**decrypt_credentials(existing_encrypted, ra.owner_id), **auth_credentials}
|
||||
stored = {} if moved else _stored_mcp_credentials(existing_doc, ra.owner_id)
|
||||
if stored:
|
||||
credentials = {**stored, **auth_credentials}
|
||||
return existing_doc, ra.owner_id, ra.access == "owner", moved, credentials
|
||||
|
||||
|
||||
@@ -168,6 +302,9 @@ class TestMCPServerConfig(Resource):
|
||||
)
|
||||
|
||||
_validate_mcp_server_url(config)
|
||||
policy_error = _mcp_policy_error(config)
|
||||
if policy_error is not None:
|
||||
return policy_error
|
||||
|
||||
ctx = _existing_mcp_context(data.get("id"), user, config)
|
||||
if not isinstance(ctx, tuple):
|
||||
@@ -183,7 +320,7 @@ class TestMCPServerConfig(Resource):
|
||||
safe_result = {
|
||||
k: v
|
||||
for k, v in result.items()
|
||||
if k in ("success", "requires_oauth", "auth_url")
|
||||
if k in ("success", "requires_oauth", "auth_url", "task_id")
|
||||
}
|
||||
return make_response(jsonify(safe_result), 200)
|
||||
|
||||
@@ -269,6 +406,9 @@ class MCPServerSave(Resource):
|
||||
)
|
||||
|
||||
_validate_mcp_server_url(config)
|
||||
policy_error = _mcp_policy_error(config)
|
||||
if policy_error is not None:
|
||||
return policy_error
|
||||
|
||||
# An existing id is always an update of THAT row, written as its
|
||||
# owner; it never falls through to creating a copy for the caller.
|
||||
@@ -280,11 +420,16 @@ class MCPServerSave(Resource):
|
||||
mcp_config = config.copy()
|
||||
mcp_config["auth_credentials"] = merged_credentials
|
||||
|
||||
if auth_type == "oauth":
|
||||
if auth_type == "oauth" and not config.get("oauth_task_id"):
|
||||
# Only the owner reaches here for an existing server (see
|
||||
# ``_existing_mcp_context``), and every OAuth save needs the
|
||||
# sign-in they just completed.
|
||||
if not config.get("oauth_task_id"):
|
||||
# ``_existing_mcp_context``). No new handshake: they signed in
|
||||
# to this server before, so its stored tokens answer the
|
||||
# discovery, and the save below still needs that connection.
|
||||
try:
|
||||
mcp_tool = MCPTool(config=mcp_config, user_id=owner_id)
|
||||
mcp_tool.discover_tools()
|
||||
actions_metadata = mcp_tool.get_actions_metadata()
|
||||
except Exception:
|
||||
return make_response(
|
||||
jsonify(
|
||||
{
|
||||
@@ -294,6 +439,7 @@ class MCPServerSave(Resource):
|
||||
),
|
||||
400,
|
||||
)
|
||||
elif auth_type == "oauth":
|
||||
redis_client = get_redis_instance()
|
||||
manager = MCPOAuthManager(redis_client)
|
||||
result = manager.get_oauth_status(
|
||||
@@ -334,9 +480,32 @@ class MCPServerSave(Resource):
|
||||
"redirect_uri",
|
||||
]:
|
||||
storage_config.pop(field, None)
|
||||
transformed_actions = transform_actions(actions_metadata)
|
||||
from docsgpt.connectors.permissions import apply_default_permissions
|
||||
|
||||
transformed_actions = apply_default_permissions(
|
||||
"mcp_tool", transform_actions(actions_metadata),
|
||||
)
|
||||
|
||||
display_name = data["displayName"]
|
||||
# The connection is the owner's: an editor's save stores the
|
||||
# secret on the owner's account, as the tool runs as its owner.
|
||||
connection_id = _mcp_connection(
|
||||
owner_id, storage_config, auth_type, merged_credentials, display_name,
|
||||
) or _previous_connection(existing_doc, storage_config, owner_id)
|
||||
if auth_type == "oauth" and not connection_id:
|
||||
# A sign-in server's tokens live on its connection. Without one
|
||||
# (it was removed, and a client cached before that answered the
|
||||
# discovery) the tool would be saved unconnected.
|
||||
return make_response(
|
||||
jsonify({
|
||||
"success": False,
|
||||
"error": "Not signed in to this server. Sign in again to connect it.",
|
||||
}),
|
||||
400,
|
||||
)
|
||||
if connection_id and auth_type != "oauth":
|
||||
# The secret lives on the connection only.
|
||||
storage_config.pop("encrypted_credentials", None)
|
||||
description = f"MCP Server: {storage_config.get('server_url', 'Unknown')}"
|
||||
status_bool = bool(data.get("status", True))
|
||||
fields_out = {
|
||||
@@ -344,7 +513,7 @@ class MCPServerSave(Resource):
|
||||
"custom_name": display_name,
|
||||
"description": description,
|
||||
"config": storage_config,
|
||||
"actions": transformed_actions,
|
||||
"connection_id": connection_id,
|
||||
}
|
||||
updated_message = (
|
||||
f"MCP server updated successfully! Discovered {len(transformed_actions)} tools."
|
||||
@@ -357,6 +526,8 @@ class MCPServerSave(Resource):
|
||||
# save doesn't flip it.
|
||||
if is_owner:
|
||||
fields_out["status"] = status_bool
|
||||
# Fixed values the owner set survive a re-save.
|
||||
fields_out["actions"] = carry_pins_between(existing_doc.get("actions"), transformed_actions)
|
||||
repo.update(str(existing_doc["id"]), owner_id, fields_out)
|
||||
saved_id = str(existing_doc["id"])
|
||||
response_data = {
|
||||
@@ -375,6 +546,9 @@ class MCPServerSave(Resource):
|
||||
(existing_by_name.get("config") or {}).get("server_url")
|
||||
== storage_config.get("server_url")
|
||||
):
|
||||
fields_out["actions"] = carry_pins_between(
|
||||
existing_by_name.get("actions"), transformed_actions,
|
||||
)
|
||||
repo.update(str(existing_by_name["id"]), user, fields_out)
|
||||
saved_id = str(existing_by_name["id"])
|
||||
response_data = {
|
||||
@@ -393,6 +567,7 @@ class MCPServerSave(Resource):
|
||||
config_requirements={},
|
||||
actions=transformed_actions,
|
||||
status=status_bool,
|
||||
connection_id=connection_id,
|
||||
)
|
||||
saved_id = str(created["id"])
|
||||
response_data = {
|
||||
@@ -458,7 +633,7 @@ class MCPOAuthCallback(Resource):
|
||||
"/api/connectors/callback-status?status=error&message=Internal+server+error:+Redis+not+available.&provider=mcp_tool"
|
||||
)
|
||||
manager = MCPOAuthManager(redis_client)
|
||||
success = manager.handle_oauth_callback(state, code, error)
|
||||
success = manager.handle_oauth_callback(state, code, error, iss=request.args.get("iss"))
|
||||
if success:
|
||||
return redirect(
|
||||
"/api/connectors/callback-status?status=success&message=Authorization+code+received+successfully.+You+can+close+this+window.&provider=mcp_tool"
|
||||
@@ -505,50 +680,36 @@ class MCPAuthStatus(Resource):
|
||||
jsonify({"success": True, "statuses": {}}), 200
|
||||
)
|
||||
|
||||
oauth_server_urls: dict = {}
|
||||
from docsgpt.connectors import service
|
||||
|
||||
# Read from connection status alone: status checks never
|
||||
# decrypt credentials.
|
||||
statuses: dict = {}
|
||||
for tool in mcp_tools:
|
||||
tool_id = str(tool["id"])
|
||||
config = tool.get("config") or {}
|
||||
auth_type = config.get("auth_type", "none")
|
||||
if auth_type == "oauth":
|
||||
server_url = config.get("server_url", "")
|
||||
if server_url:
|
||||
parsed = urlparse(server_url)
|
||||
base_url = f"{parsed.scheme}://{parsed.netloc}"
|
||||
oauth_server_urls[tool_id] = (tool.get("user_id") or user, base_url)
|
||||
row = None
|
||||
if tool.get("connection_id"):
|
||||
row = sessions_repo.get(str(tool["connection_id"]))
|
||||
elif auth_type == "oauth" and config.get("server_url"):
|
||||
parsed = urlparse(config["server_url"])
|
||||
# A team-shared server signs in as its owner.
|
||||
row = sessions_repo.get_by_user_provider(
|
||||
tool.get("user_id") or user,
|
||||
service.mcp_provider(f"{parsed.scheme}://{parsed.netloc}"),
|
||||
)
|
||||
if row is not None:
|
||||
connected = service.normalize_status(row) == service.STATUS_CONNECTED
|
||||
if auth_type == "oauth" or not connected:
|
||||
statuses[tool_id] = "connected" if connected else "needs_auth"
|
||||
else:
|
||||
statuses[tool_id] = "needs_auth"
|
||||
statuses[tool_id] = "configured"
|
||||
elif auth_type == "oauth":
|
||||
statuses[tool_id] = "needs_auth"
|
||||
else:
|
||||
statuses[tool_id] = "configured"
|
||||
|
||||
if oauth_server_urls:
|
||||
# Look up a session per distinct base URL. MCP sessions
|
||||
# are stored with ``provider = "mcp:<server_url>"``
|
||||
# and the URL in ``server_url``; reuse the repo's
|
||||
# per-URL accessor rather than an ad-hoc $in query.
|
||||
url_has_tokens: dict = {}
|
||||
for owner_id, base_url in set(oauth_server_urls.values()):
|
||||
session = sessions_repo.get_by_user_and_server_url(
|
||||
owner_id, base_url,
|
||||
)
|
||||
tokens = (
|
||||
(session or {}).get("session_data", {}) or {}
|
||||
).get("tokens", {}) or {}
|
||||
# MCP code also stashes tokens into token_info on
|
||||
# the row; consider either present as "connected".
|
||||
token_info = (session or {}).get("token_info") or {}
|
||||
url_has_tokens[(owner_id, base_url)] = bool(
|
||||
tokens.get("access_token")
|
||||
or token_info.get("access_token")
|
||||
)
|
||||
|
||||
for tool_id, key in oauth_server_urls.items():
|
||||
if url_has_tokens.get(key):
|
||||
statuses[tool_id] = "connected"
|
||||
else:
|
||||
statuses[tool_id] = "needs_auth"
|
||||
|
||||
return make_response(jsonify({"success": True, "statuses": statuses}), 200)
|
||||
except Exception as e:
|
||||
current_app.logger.error(
|
||||
|
||||
@@ -19,6 +19,7 @@ from docsgpt.agents.default_tools import (
|
||||
WORKFLOW_ONLY_BUILTINS,
|
||||
)
|
||||
from docsgpt.agents.tool_executor import API_TOOL_SECRET_SECTIONS, API_TOOL_SECRETS_KEY
|
||||
from docsgpt.agents.tool_pins import iter_parameters, llm_fills, merge_submitted_actions, PinChangeRefused
|
||||
from docsgpt.agents.tools.spec_parser import parse_spec
|
||||
from docsgpt.agents.tools.tool_manager import ToolManager
|
||||
from docsgpt.api import api
|
||||
@@ -33,9 +34,13 @@ from docsgpt.api.user.resource_access import (
|
||||
settings_many,
|
||||
)
|
||||
from docsgpt.api.user.team_sharing import visible_with_access
|
||||
from docsgpt.connectors.catalog import base_url, definition_for_tool
|
||||
from docsgpt.connectors.resolve import carry_removed_connection
|
||||
from docsgpt.connectors.service import account_tool_names
|
||||
from docsgpt.connectors.permissions import owner_credential_writes
|
||||
from docsgpt.core.settings import settings
|
||||
from docsgpt.core.url_validation import SSRFError, validate_url
|
||||
from docsgpt.security.encryption import decrypt_credentials, encrypt_credentials
|
||||
from docsgpt.security.encryption import CredentialDecryptionError, decrypt_credentials, encrypt_credentials
|
||||
from docsgpt.storage.db.base_repository import looks_like_uuid
|
||||
from docsgpt.storage.db.repositories.artifacts import ArtifactsRepository
|
||||
from docsgpt.storage.db.repositories.notes import NotesRepository
|
||||
@@ -189,7 +194,14 @@ def _merge_secrets_on_update(new_config, existing_config, config_requirements, u
|
||||
_CREDENTIALS_FOR_NEW_SERVER = "Enter credentials for the new server"
|
||||
_FORBIDDEN_MESSAGE = "Your access to this item doesn't allow that"
|
||||
_MCP_CREDENTIAL_AUTH_TYPES = {"api_key", "bearer", "basic"}
|
||||
_META_KEYS = ("name", "displayName", "customName", "description", "actions")
|
||||
# ``name`` (the tool type) and ``actions`` are handled on their own.
|
||||
_META_KEYS = ("displayName", "customName", "description")
|
||||
_TYPE_IS_FIXED = "A tool's type can't be changed"
|
||||
_FIXED_VALUES_OWNER_ONLY = "Only the tool's owner can change fixed values"
|
||||
_MOVE_THROUGH_MCP_SAVE = (
|
||||
"This server signs in through a connection: change its address or sign-in by saving the "
|
||||
"server again (/api/mcp_server/save)"
|
||||
)
|
||||
|
||||
|
||||
class CredentialsRequired(Exception):
|
||||
@@ -359,6 +371,96 @@ def _api_tool_config_needs_credentials(new_config: dict, existing_config: dict)
|
||||
return False
|
||||
|
||||
|
||||
def _api_tool_param_state(section: str, spec: dict, *, stored: bool) -> tuple:
|
||||
"""Who fills an api_tool parameter and which value it keeps.
|
||||
|
||||
A header / query value is only ever shown masked, so a stored one (sealed,
|
||||
or legacy plaintext) reads as ``"stored"`` and any value a client sends
|
||||
is a new one.
|
||||
|
||||
Args:
|
||||
section: ``headers``, ``query_params`` or ``body``.
|
||||
spec: The parameter's schema.
|
||||
stored: Whether ``spec`` comes from the stored config.
|
||||
|
||||
Returns:
|
||||
``(filled_by_llm, value marker)``; the marker is None without a value.
|
||||
"""
|
||||
value = spec.get("value")
|
||||
if section in API_TOOL_SECRET_SECTIONS:
|
||||
if _has_value(value):
|
||||
marker: Any = "stored" if stored else ("new", value)
|
||||
else:
|
||||
marker = "stored" if spec.get("has_value") else None
|
||||
else:
|
||||
marker = value if _has_value(value) else None
|
||||
return llm_fills(spec), marker
|
||||
|
||||
|
||||
def _api_tool_fixed_values_changed(new_config: dict, existing_config: dict) -> bool:
|
||||
"""Whether an api_tool save changes a fixed value of an existing action.
|
||||
|
||||
A fixed value may be a secret (an API key in a query) or where a call
|
||||
goes, so changing who fills a parameter, or its value, or clearing a
|
||||
stored one is the owner's. Flipping ``filled_by_llm`` on a stored secret
|
||||
would hand it to the model and show it in the chat. A parameter that
|
||||
carries no value can be added or removed freely; a new action brings a
|
||||
URL and is judged by :func:`_api_tool_config_needs_credentials`.
|
||||
|
||||
Args:
|
||||
new_config: The config the client sent.
|
||||
existing_config: The stored config.
|
||||
|
||||
Returns:
|
||||
True when any existing action's fixed values would differ.
|
||||
"""
|
||||
old_actions = (existing_config or {}).get("actions") or {}
|
||||
new_actions = (new_config or {}).get("actions") or {}
|
||||
if not isinstance(old_actions, dict) or not isinstance(new_actions, dict):
|
||||
return bool(old_actions) or bool(new_actions)
|
||||
for name, action in new_actions.items():
|
||||
old = old_actions.get(name)
|
||||
if not isinstance(old, dict) or not isinstance(action, dict):
|
||||
continue
|
||||
before = {(s, p): _api_tool_param_state(s, d, stored=True) for s, p, d in iter_parameters(old)}
|
||||
after = {(s, p): _api_tool_param_state(s, d, stored=False) for s, p, d in iter_parameters(action)}
|
||||
for key in before.keys() | after.keys():
|
||||
if key in before and key in after:
|
||||
if before[key] != after[key]:
|
||||
return True
|
||||
elif (before.get(key) or after.get(key))[1] is not None:
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
def _connection_server_moved(tool_doc: dict, new_config: Optional[dict]) -> bool:
|
||||
"""Whether a config save re-points a connection-backed MCP tool.
|
||||
|
||||
The connection holds the key for one server and one way of signing in;
|
||||
a tool moved on these routes would keep a connection that no longer
|
||||
applies (it is refused at run time) and the new key would be stored
|
||||
where nothing reads it. ``/api/mcp_server/save`` moves a server properly.
|
||||
|
||||
Args:
|
||||
tool_doc: The stored ``user_tools`` row.
|
||||
new_config: The incoming ``config``.
|
||||
|
||||
Returns:
|
||||
True when the base URL or ``auth_type`` would change.
|
||||
"""
|
||||
if tool_doc.get("name") != "mcp_tool" or not tool_doc.get("connection_id"):
|
||||
return False
|
||||
existing = tool_doc.get("config") or {}
|
||||
new_config = new_config if isinstance(new_config, dict) else {}
|
||||
if "server_url" in new_config and base_url(str(new_config.get("server_url") or "").strip()) != base_url(
|
||||
existing.get("server_url")
|
||||
):
|
||||
return True
|
||||
return "auth_type" in new_config and (new_config.get("auth_type") or "none") != (
|
||||
existing.get("auth_type") or "none"
|
||||
)
|
||||
|
||||
|
||||
def _mcp_origin_changed(new_config: dict, existing_config: dict) -> bool:
|
||||
"""Whether a save moves an MCP server to another scheme, host or port."""
|
||||
old_url = (existing_config or {}).get("server_url")
|
||||
@@ -396,18 +498,40 @@ def check_oauth_mcp_owner_only(
|
||||
raise AccessDenied(403, SHARED_OAUTH_OWNER_ONLY)
|
||||
|
||||
|
||||
def check_api_tool_fixed_values(
|
||||
ra: ResourceAccess, new_config: Optional[dict], existing_config: Optional[dict]
|
||||
) -> None:
|
||||
"""Keep an api_tool's fixed header, query and body values with its owner.
|
||||
|
||||
Args:
|
||||
ra: The caller's access to the tool.
|
||||
new_config: The incoming ``config``.
|
||||
existing_config: The stored ``config``.
|
||||
|
||||
Raises:
|
||||
AccessDenied: 403 when a non-owner changes a fixed value (see
|
||||
:func:`_api_tool_fixed_values_changed`).
|
||||
"""
|
||||
if ra.access != "owner" and _api_tool_fixed_values_changed(new_config or {}, existing_config or {}):
|
||||
raise AccessDenied(403, _FIXED_VALUES_OWNER_ONLY)
|
||||
|
||||
|
||||
def _prepare_tool_config(tool_doc: dict, new_config: dict, config_requirements: dict) -> dict:
|
||||
"""Validate-free merge of an incoming config with the stored one, as the owner.
|
||||
|
||||
Handles the three secret stores: ``config_requirements`` secrets
|
||||
(``encrypted_credentials``), api_tool header/query values, and the MCP
|
||||
origin-change rule (a new scheme, host or port drops stored credentials).
|
||||
A removed connection's note is carried over from the stored config; the
|
||||
client's copy is ignored.
|
||||
|
||||
Raises:
|
||||
CredentialsRequired: the MCP origin changed and no new secret arrived.
|
||||
"""
|
||||
owner_id = tool_doc["user_id"]
|
||||
existing_config = tool_doc.get("config") or {}
|
||||
# The note that its connection was removed is the server's to keep.
|
||||
new_config = carry_removed_connection(new_config, existing_config)
|
||||
if tool_doc.get("name") == "api_tool":
|
||||
return _seal_api_tool_secrets(new_config, existing_config, owner_id)
|
||||
moved = tool_doc.get("name") == "mcp_tool" and _mcp_origin_changed(new_config, existing_config)
|
||||
@@ -479,8 +603,30 @@ def transform_actions(actions_metadata):
|
||||
return transformed
|
||||
|
||||
|
||||
def _stored_actions(tool_doc: dict) -> list:
|
||||
"""A tool's stored actions, or its class's own for a row stored without any.
|
||||
|
||||
Args:
|
||||
tool_doc: The ``user_tools`` row.
|
||||
|
||||
Returns:
|
||||
The action list submitted actions are validated against.
|
||||
"""
|
||||
actions = tool_doc.get("actions") or []
|
||||
if actions or tool_doc.get("name") in ("mcp_tool", "api_tool"):
|
||||
return actions
|
||||
tool_instance = tool_manager.tools.get(tool_doc.get("name"))
|
||||
if tool_instance is None:
|
||||
return actions
|
||||
return transform_actions(copy.deepcopy(tool_instance.get_actions_metadata()))
|
||||
|
||||
|
||||
tools_ns = Namespace("tools", description="Tool management operations", path="/api")
|
||||
|
||||
# Tools the Connectors page adds through "Add custom connector" rather than
|
||||
# the Add Tool modal.
|
||||
_CUSTOM_CONNECTOR_TOOLS = {"mcp_tool": "custom_mcp", "api_tool": "custom_openapi"}
|
||||
|
||||
|
||||
@tools_ns.route("/available_tools")
|
||||
class AvailableTools(Resource):
|
||||
@@ -488,6 +634,16 @@ class AvailableTools(Resource):
|
||||
def get(self):
|
||||
if not request.decoded_token:
|
||||
return make_response(jsonify({"success": False}), 401)
|
||||
from docsgpt.connectors import service as connection_service
|
||||
|
||||
try:
|
||||
with db_readonly() as conn:
|
||||
policies = connection_service.load_policies(conn)
|
||||
except Exception:
|
||||
# Without the admin's switches, fall back to the defaults (a
|
||||
# connector is on when its server settings are present).
|
||||
current_app.logger.warning("Could not read connector policies", exc_info=True)
|
||||
policies = {}
|
||||
try:
|
||||
tools_metadata = []
|
||||
for tool_name, tool_instance in tool_manager.tools.items():
|
||||
@@ -497,6 +653,19 @@ class AvailableTools(Resource):
|
||||
description = lines[1].strip() if len(lines) > 1 else ""
|
||||
config_req = tool_instance.get_config_requirements()
|
||||
actions = tool_instance.get_actions_metadata()
|
||||
definition = definition_for_tool(tool_name)
|
||||
if definition is not None:
|
||||
if not (definition.configured and connection_service.connector_is_enabled(
|
||||
policies, definition.key,
|
||||
)):
|
||||
continue
|
||||
group, connector_key = "service", definition.key
|
||||
# One name everywhere: the connector's, not the tool's own.
|
||||
name = definition.name
|
||||
elif tool_name in _CUSTOM_CONNECTOR_TOOLS:
|
||||
group, connector_key = "custom", _CUSTOM_CONNECTOR_TOOLS[tool_name]
|
||||
else:
|
||||
group, connector_key = "built_in", None
|
||||
tools_metadata.append(
|
||||
{
|
||||
"name": tool_name,
|
||||
@@ -504,6 +673,8 @@ class AvailableTools(Resource):
|
||||
"description": description,
|
||||
"configRequirements": config_req,
|
||||
"actions": actions,
|
||||
"group": group,
|
||||
"connector_key": connector_key,
|
||||
}
|
||||
)
|
||||
except Exception as err:
|
||||
@@ -534,6 +705,8 @@ class GetTools(Resource):
|
||||
team_shared = visible_with_access(conn, user, "tool")
|
||||
shared_ids = [tid for tid in team_shared if tid not in owned_ids]
|
||||
shared_rows = tools_repo.list_by_ids(shared_ids)
|
||||
# "Telegram · Alerts bot" when the owner has several bots.
|
||||
account_names = account_tool_names(conn, [*rows, *shared_rows])
|
||||
switches = settings_many(conn, "tool", [*owned_ids, *shared_ids])
|
||||
prefs = UserToolPreferencesRepository(conn).in_chat_many(user, shared_ids)
|
||||
shared_via = _shared_via(conn, user, shared_ids)
|
||||
@@ -542,6 +715,9 @@ class GetTools(Resource):
|
||||
|
||||
def _shape_tool(row, *, ownership="user", force_strip_secret=False):
|
||||
tool_copy = _row_to_api(row)
|
||||
# The writes an agent's API write allowlist can cover (read
|
||||
# from the stored row, before any secret is masked).
|
||||
tool_copy["owner_credential_writes"] = owner_credential_writes(row)
|
||||
config_req = tool_copy.get("configRequirements", {})
|
||||
if not config_req:
|
||||
tool_instance = tool_manager.tools.get(tool_copy.get("name"))
|
||||
@@ -556,10 +732,16 @@ class GetTools(Resource):
|
||||
):
|
||||
tool_copy["config"]["has_encrypted_credentials"] = True
|
||||
tool_copy["config"].pop("encrypted_credentials", None)
|
||||
if tool_copy.get("connection_id"):
|
||||
# The secret lives on the connection; the form must not
|
||||
# ask for it again.
|
||||
tool_copy.setdefault("config", {})["has_encrypted_credentials"] = True
|
||||
if tool_copy.get("name") == "api_tool":
|
||||
# Header / query-param values are secrets for everyone.
|
||||
tool_copy["config"] = mask_api_tool_config(tool_copy.get("config") or {})
|
||||
tool_copy["ownership"] = ownership
|
||||
if str(row["id"]) in account_names:
|
||||
tool_copy["customName"] = tool_copy["displayName"] = account_names[str(row["id"])]
|
||||
return tool_copy
|
||||
|
||||
for row in rows:
|
||||
@@ -657,6 +839,9 @@ class CreateTool(Resource):
|
||||
missing_fields = check_required_fields(data, required_fields)
|
||||
if missing_fields:
|
||||
return missing_fields
|
||||
if isinstance(data.get("config"), dict):
|
||||
# Only removing a connection notes that it was removed.
|
||||
data["config"] = carry_removed_connection(data["config"], None)
|
||||
try:
|
||||
if data["name"] == "mcp_tool":
|
||||
server_url = (data.get("config", {}).get("server_url") or "").strip()
|
||||
@@ -680,6 +865,11 @@ class CreateTool(Resource):
|
||||
f"Error getting tool actions: {err}", exc_info=True
|
||||
)
|
||||
return make_response(jsonify({"success": False}), 400)
|
||||
definition = definition_for_tool(data["name"])
|
||||
if definition is not None:
|
||||
connected = _create_connected_tool(user, data, definition, tool_instance)
|
||||
if connected is not None:
|
||||
return connected
|
||||
try:
|
||||
config_requirements = tool_instance.get_config_requirements()
|
||||
if config_requirements:
|
||||
@@ -722,6 +912,82 @@ class CreateTool(Resource):
|
||||
return make_response(jsonify({"id": new_id}), 200)
|
||||
|
||||
|
||||
def _create_connected_tool(user, data, definition, tool_instance):
|
||||
"""Create a service tool whose secret lives on a connection, not the tool.
|
||||
|
||||
Uses ``connection_id`` when given, otherwise stores the pasted secret on
|
||||
a connection (reusing an identical one). Returns None to fall back to the
|
||||
legacy path when a multi-user install still runs on the default key.
|
||||
"""
|
||||
from docsgpt.connectors import catalog as connector_catalog
|
||||
from docsgpt.connectors import service as connection_service
|
||||
from docsgpt.storage.db.repositories.connector_sessions import ConnectorSessionsRepository
|
||||
|
||||
config_requirements = tool_instance.get_config_requirements()
|
||||
public, secrets = connection_service.split_secrets(data.get("config") or {}, config_requirements)
|
||||
connection_id = data.get("connection_id")
|
||||
# An existing connection supplies the secrets the request leaves out.
|
||||
validation_errors = _validate_config(
|
||||
data.get("config") or {}, config_requirements, has_existing_secrets=bool(connection_id),
|
||||
)
|
||||
if validation_errors:
|
||||
return make_response(
|
||||
jsonify({"success": False, "message": "Validation failed", "errors": validation_errors}), 400,
|
||||
)
|
||||
try:
|
||||
with db_session() as conn:
|
||||
if connection_id:
|
||||
connection = ConnectorSessionsRepository(conn).get_for_user(str(connection_id), user)
|
||||
if connection is None or connector_catalog.connector_key_for_row(connection) != definition.key:
|
||||
return make_response(jsonify({"success": False, "message": "Connection not found"}), 404)
|
||||
else:
|
||||
connection, _ = connection_service.create_api_key_connection(conn, user, definition, secrets)
|
||||
created = connection_service.create_tool_for_connection(
|
||||
conn,
|
||||
user,
|
||||
connection,
|
||||
template=data["name"],
|
||||
display_name=data.get("customName") or data.get("displayName") or definition.name,
|
||||
config=public,
|
||||
status=bool(data.get("status", True)),
|
||||
)
|
||||
except connection_service.EncryptionKeyNotConfigured:
|
||||
return None
|
||||
except connection_service.ConnectorDisabled as err:
|
||||
return make_response(jsonify({"success": False, "message": str(err)}), 403)
|
||||
except ValueError as err:
|
||||
return make_response(jsonify({"success": False, "message": str(err)}), 400)
|
||||
return make_response(jsonify({"id": str(created["id"]), "connection_id": str(connection["id"])}), 200)
|
||||
|
||||
|
||||
def _update_connection_secrets(conn, user, tool_doc, config, config_requirements):
|
||||
"""Write changed secrets onto the tool's connection; return the tool's public config.
|
||||
|
||||
Returns None when the caller does not own the connection (an editor on a
|
||||
team share may change actions, never the owner's credentials).
|
||||
"""
|
||||
from docsgpt.connectors import service as connection_service
|
||||
from docsgpt.storage.db.repositories.connector_sessions import ConnectorSessionsRepository
|
||||
|
||||
public, secrets = connection_service.split_secrets(config or {}, config_requirements)
|
||||
if secrets:
|
||||
connection = ConnectorSessionsRepository(conn).get_for_user(str(tool_doc["connection_id"]), user)
|
||||
if connection is None:
|
||||
return None
|
||||
try:
|
||||
stored = connection_service.read_secrets(connection)
|
||||
except CredentialDecryptionError:
|
||||
# Unreadable after a lost key: the new secret replaces it.
|
||||
stored = {}
|
||||
credentials = {**(stored.get("credentials") or {}), **secrets}
|
||||
connection_service.write_secrets(
|
||||
conn, connection, {**stored, "credentials": credentials},
|
||||
status=connection_service.STATUS_CONNECTED, last_error=None,
|
||||
)
|
||||
connection_service.resume_sources(conn, str(connection["id"]))
|
||||
return public
|
||||
|
||||
|
||||
@tools_ns.route("/update_tool")
|
||||
class UpdateTool(Resource):
|
||||
@api.expect(
|
||||
@@ -743,6 +1009,16 @@ class UpdateTool(Resource):
|
||||
)
|
||||
@api.doc(description="Update a tool by ID")
|
||||
def post(self):
|
||||
"""Update a tool's names, actions, config or chat switch.
|
||||
|
||||
The tool type (``name``) never changes. ``actions`` are checked
|
||||
against the stored ones like ``/api/update_tool_actions``: nothing
|
||||
is added and fixed values are the owner's. A connection-backed MCP
|
||||
server is moved only through ``/api/mcp_server/save``.
|
||||
|
||||
Returns:
|
||||
``{"success": true}``, or 400 / 403 / 404 with a message.
|
||||
"""
|
||||
decoded_token = request.decoded_token
|
||||
if not decoded_token:
|
||||
return make_response(jsonify({"success": False}), 401)
|
||||
@@ -817,13 +1093,32 @@ class UpdateTool(Resource):
|
||||
return make_response(
|
||||
jsonify({"success": False, "message": "Tool not found"}), 404,
|
||||
)
|
||||
if update_data:
|
||||
if "name" in data and data["name"] != tool_doc.get("name"):
|
||||
# The type decides what the config and the connection's
|
||||
# credentials are used for; it is set once, on create.
|
||||
return make_response(jsonify({"success": False, "message": _TYPE_IS_FIXED}), 400)
|
||||
if update_data or "actions" in data:
|
||||
check_action(ra, "edit")
|
||||
if "actions" in data:
|
||||
# The stored actions are the schema, and a fixed value is
|
||||
# the owner's (as on /api/update_tool_actions).
|
||||
try:
|
||||
update_data["actions"] = merge_submitted_actions(
|
||||
_stored_actions(tool_doc), data["actions"], may_change_pins=ra.access == "owner",
|
||||
)
|
||||
except PinChangeRefused as err:
|
||||
return make_response(jsonify({"success": False, "message": str(err)}), 403)
|
||||
except ValueError as err:
|
||||
return make_response(jsonify({"success": False, "message": str(err)}), 400)
|
||||
if "config" in data:
|
||||
tool_name = tool_doc.get("name", data.get("name"))
|
||||
tool_name = tool_doc.get("name")
|
||||
existing_config = tool_doc.get("config", {}) or {}
|
||||
if tool_name == "mcp_tool":
|
||||
check_oauth_mcp_owner_only(ra, existing_config, data["config"])
|
||||
if _connection_server_moved(tool_doc, data["config"]):
|
||||
return make_response(jsonify({"success": False, "message": _MOVE_THROUGH_MCP_SAVE}), 400)
|
||||
if tool_name == "api_tool":
|
||||
check_api_tool_fixed_values(ra, data["config"], existing_config)
|
||||
if tool_name == "api_tool" and not _api_tool_config_needs_credentials(
|
||||
data["config"], existing_config
|
||||
):
|
||||
@@ -834,7 +1129,11 @@ class UpdateTool(Resource):
|
||||
config_requirements = (
|
||||
tool_instance.get_config_requirements() if tool_instance else {}
|
||||
)
|
||||
has_existing_secrets = "encrypted_credentials" in existing_config
|
||||
has_existing_secrets = (
|
||||
"encrypted_credentials" in existing_config or bool(tool_doc.get("connection_id"))
|
||||
)
|
||||
# Validate before touching the connection: a rejected
|
||||
# edit must not have rotated its credentials.
|
||||
if config_requirements:
|
||||
validation_errors = _validate_config(
|
||||
data["config"], config_requirements,
|
||||
@@ -849,6 +1148,20 @@ class UpdateTool(Resource):
|
||||
}),
|
||||
400,
|
||||
)
|
||||
if tool_doc.get("connection_id"):
|
||||
# The connection may back the owner's other tools too,
|
||||
# so its secret stays the owner's to change even when
|
||||
# an editor may change the tool's own credentials.
|
||||
new_config = _update_connection_secrets(
|
||||
conn, user, tool_doc, data["config"], config_requirements,
|
||||
)
|
||||
if new_config is None:
|
||||
return make_response(
|
||||
jsonify({"success": False, "message": "Only the owner can change the credentials"}),
|
||||
403,
|
||||
)
|
||||
data = {**data, "config": new_config}
|
||||
|
||||
update_data["config"] = _prepare_tool_config(
|
||||
tool_doc, data["config"], config_requirements
|
||||
)
|
||||
@@ -891,6 +1204,15 @@ class UpdateToolConfig(Resource):
|
||||
)
|
||||
@api.doc(description="Update the configuration of a tool")
|
||||
def post(self):
|
||||
"""Replace a tool's config, keeping stored secrets the client left out.
|
||||
|
||||
A connection-backed tool's new key goes to its connection (the
|
||||
owner's to change) and its server cannot be moved here; an api_tool's
|
||||
fixed values are the owner's.
|
||||
|
||||
Returns:
|
||||
``{"success": true}``, or 400 / 403 / 404 with a message.
|
||||
"""
|
||||
decoded_token = request.decoded_token
|
||||
if not decoded_token:
|
||||
return make_response(jsonify({"success": False}), 401)
|
||||
@@ -931,12 +1253,18 @@ class UpdateToolConfig(Resource):
|
||||
jsonify({"success": False, "message": "Invalid server URL"}),
|
||||
400,
|
||||
)
|
||||
if _connection_server_moved(tool_doc, data["config"]):
|
||||
return make_response(jsonify({"success": False, "message": _MOVE_THROUGH_MCP_SAVE}), 400)
|
||||
if tool_name == "api_tool":
|
||||
check_api_tool_fixed_values(ra, data["config"], tool_doc.get("config"))
|
||||
tool_instance = tool_manager.tools.get(tool_name)
|
||||
config_requirements = (
|
||||
tool_instance.get_config_requirements() if tool_instance else {}
|
||||
)
|
||||
existing_config = tool_doc.get("config", {}) or {}
|
||||
has_existing_secrets = "encrypted_credentials" in existing_config
|
||||
has_existing_secrets = (
|
||||
"encrypted_credentials" in existing_config or bool(tool_doc.get("connection_id"))
|
||||
)
|
||||
|
||||
if config_requirements:
|
||||
validation_errors = _validate_config(
|
||||
@@ -953,7 +1281,17 @@ class UpdateToolConfig(Resource):
|
||||
400,
|
||||
)
|
||||
|
||||
final_config = _prepare_tool_config(tool_doc, data["config"], config_requirements)
|
||||
config = data["config"]
|
||||
if tool_doc.get("connection_id"):
|
||||
# The tool runs with its connection's key: a new one goes
|
||||
# there (the owner's to change), not into this config.
|
||||
config = _update_connection_secrets(conn, user, tool_doc, config, config_requirements)
|
||||
if config is None:
|
||||
return make_response(
|
||||
jsonify({"success": False, "message": "Only the owner can change the credentials"}),
|
||||
403,
|
||||
)
|
||||
final_config = _prepare_tool_config(tool_doc, config, config_requirements)
|
||||
|
||||
repo.update(str(tool_doc["id"]), ra.owner_id, {"config": final_config})
|
||||
except AccessDenied as err:
|
||||
@@ -1009,7 +1347,23 @@ class UpdateToolActions(Resource):
|
||||
# ``edit`` covers action on/off, descriptions and approval
|
||||
# (``require_approval``); actions carry no credentials.
|
||||
ra = require(conn, "tool", data["id"], user, "edit")
|
||||
UserToolsRepository(conn).update(ra.resource_id, ra.owner_id, {"actions": data["actions"]})
|
||||
tool_doc = _load_owned_row(conn, ra)
|
||||
if not tool_doc:
|
||||
return make_response(
|
||||
jsonify({"success": False, "message": "Tool not found"}), 404,
|
||||
)
|
||||
# The stored actions are the schema: nothing can be added,
|
||||
# and a fixed value is the owner's to set (an editor pinning
|
||||
# a Telegram chat would redirect the owner's bot).
|
||||
try:
|
||||
actions = merge_submitted_actions(
|
||||
_stored_actions(tool_doc), data["actions"], may_change_pins=ra.access == "owner",
|
||||
)
|
||||
except PinChangeRefused as err:
|
||||
return make_response(jsonify({"success": False, "message": str(err)}), 403)
|
||||
except ValueError as err:
|
||||
return make_response(jsonify({"success": False, "message": str(err)}), 400)
|
||||
UserToolsRepository(conn).update(str(tool_doc["id"]), ra.owner_id, {"actions": actions})
|
||||
except AccessDenied as err:
|
||||
return denied_response(err)
|
||||
except Exception as err:
|
||||
|
||||
@@ -10,12 +10,20 @@ from docsgpt.agents.workflows.cel_evaluator import (
|
||||
CelEvaluationError,
|
||||
validate_cel_expression,
|
||||
)
|
||||
from docsgpt.api.user import resource_access
|
||||
from docsgpt.api.user.resource_access import (
|
||||
AccessDenied,
|
||||
best_effort,
|
||||
cached_resolves,
|
||||
can_use_ref,
|
||||
named_ref_keys,
|
||||
parse_confirmations,
|
||||
plan_sponsors,
|
||||
resolve,
|
||||
resource_states,
|
||||
sponsor_audience,
|
||||
sponsor_details,
|
||||
sponsors_after_save,
|
||||
sponsor_refusal,
|
||||
)
|
||||
from docsgpt.storage.db.base_repository import looks_like_uuid
|
||||
from docsgpt.storage.db.repositories.workflow_edges import WorkflowEdgesRepository
|
||||
@@ -125,6 +133,44 @@ def _node_refs(nodes: List[Dict]) -> List[Tuple[str, str]]:
|
||||
return refs
|
||||
|
||||
|
||||
def _node_ref_details(nodes: List[Dict], visible: Optional[Set[str]] = None) -> Dict[str, List[Dict]]:
|
||||
"""Names of the tools and sources the graph's agent nodes reference.
|
||||
|
||||
Looked up by id whoever owns them, so an editor's node pickers can show
|
||||
(and remove) the owner's private tools and sources. Builtin tool ids
|
||||
always resolve; any other id only when its ``"<type>:<id>"`` key is in
|
||||
``visible`` (see ``resource_access.named_ref_keys``), so a node naming
|
||||
someone else's resource never reveals its name.
|
||||
|
||||
Args:
|
||||
nodes: Nodes in builder shape.
|
||||
visible: Keys whose names may be read; None allows every id.
|
||||
|
||||
Returns:
|
||||
dict: ``tools`` as ``[{id, name, display_name}]`` and ``sources`` as
|
||||
``[{id, name}]``, each id once.
|
||||
"""
|
||||
from docsgpt.agents.default_tools import is_synthesized_tool_id
|
||||
from docsgpt.api.user.base import resolve_source_details, resolve_tool_details
|
||||
|
||||
tool_ids: List[str] = []
|
||||
source_ids: List[str] = []
|
||||
for resource_type, resource_id in _node_refs(nodes):
|
||||
if (
|
||||
visible is not None
|
||||
and not (resource_type == "tool" and is_synthesized_tool_id(resource_id))
|
||||
and f"{resource_type}:{resource_id.lower()}" not in visible
|
||||
):
|
||||
continue
|
||||
bucket = tool_ids if resource_type == "tool" else source_ids
|
||||
if resource_id not in bucket:
|
||||
bucket.append(resource_id)
|
||||
return {
|
||||
"tools": resolve_tool_details(tool_ids),
|
||||
"sources": resolve_source_details(source_ids),
|
||||
}
|
||||
|
||||
|
||||
def _new_node_ref_denied(
|
||||
conn, previous_nodes: List[Dict], new_nodes: List[Dict], caller: str
|
||||
) -> Optional[AccessDenied]:
|
||||
@@ -133,13 +179,16 @@ def _new_node_ref_denied(
|
||||
A workflow runs as its owner, so an editor saving the owner's graph must
|
||||
not reference the owner's private tools or sources: the caller's own
|
||||
access counts (``use_in_own`` for a tool, ``use`` for a source), not the
|
||||
owner's. Refs already in the stored graph stay, like an agent's.
|
||||
owner's. The owner's own saves are checked the same way, so no graph
|
||||
names a resource its owner never could use. Refs already in the stored
|
||||
graph stay, like an agent's.
|
||||
|
||||
Args:
|
||||
conn: Open database connection.
|
||||
previous_nodes: The stored graph's nodes, in builder shape.
|
||||
previous_nodes: The stored graph's nodes, in builder shape (empty
|
||||
when creating the workflow).
|
||||
new_nodes: The nodes being saved.
|
||||
caller: The editor saving.
|
||||
caller: The user saving, owner or editor.
|
||||
|
||||
Returns:
|
||||
An :class:`AccessDenied` to return, or None when every new ref is fine.
|
||||
@@ -600,6 +649,9 @@ class WorkflowList(Resource):
|
||||
|
||||
try:
|
||||
with db_session() as conn:
|
||||
denied = _new_node_ref_denied(conn, [], nodes_data, user_id)
|
||||
if denied is not None:
|
||||
return _denied(denied)
|
||||
repo = WorkflowsRepository(conn)
|
||||
workflow = repo.create(user_id, name, description=description)
|
||||
pg_workflow_id = str(workflow["id"])
|
||||
@@ -631,16 +683,45 @@ class WorkflowDetail(Resource):
|
||||
edges = WorkflowEdgesRepository(conn).find_by_version(
|
||||
pg_workflow_id, graph_version,
|
||||
)
|
||||
sponsored = sponsor_details(conn, "workflow", workflow)
|
||||
serialized_nodes = [serialize_node(n) for n in nodes]
|
||||
# Edit-page detail (sponsors, run state, names of node
|
||||
# resources) only for people who may edit the workflow.
|
||||
sponsored: list = []
|
||||
states: list = []
|
||||
audience = None
|
||||
visible: Optional[Set[str]] = None
|
||||
with cached_resolves():
|
||||
if resource_access.holder_editable_by(conn, "workflow", workflow, user_id):
|
||||
refs = _node_refs(serialized_nodes)
|
||||
sponsored = sponsor_details(conn, "workflow", workflow, viewer=user_id)
|
||||
# Run state never fails the read; node names follow
|
||||
# the same rule as the state's names.
|
||||
states = best_effort(
|
||||
conn, "the workflow's resource states",
|
||||
lambda: resource_states(conn, "workflow", workflow, refs, user_id), [],
|
||||
)
|
||||
audience = best_effort(
|
||||
conn, "the workflow's audience",
|
||||
lambda: sponsor_audience(conn, "workflow", workflow, states, sponsored), None,
|
||||
)
|
||||
visible = named_ref_keys(states)
|
||||
ref_details = (
|
||||
_node_ref_details(serialized_nodes, visible)
|
||||
if visible is not None
|
||||
else {"tools": [], "sources": []}
|
||||
)
|
||||
except Exception as err:
|
||||
return _workflow_error_response("Failed to fetch workflow", err)
|
||||
|
||||
return success_response(
|
||||
{
|
||||
"workflow": serialize_workflow(workflow),
|
||||
"nodes": [serialize_node(n) for n in nodes],
|
||||
"nodes": serialized_nodes,
|
||||
"edges": [serialize_edge(e) for e in edges],
|
||||
"resource_sponsors": sponsored,
|
||||
"resource_states": states,
|
||||
"ref_details": ref_details,
|
||||
**({"sponsor_audience": audience} if audience is not None else {}),
|
||||
}
|
||||
)
|
||||
|
||||
@@ -674,16 +755,35 @@ class WorkflowDetail(Resource):
|
||||
nodes_data = normalize_agent_node_json_schemas(nodes_data)
|
||||
pg_workflow_id = str(workflow["id"])
|
||||
current_graph_version = get_workflow_graph_version(workflow)
|
||||
if acting != user_id:
|
||||
previous_nodes = [
|
||||
serialize_node(n)
|
||||
for n in WorkflowNodesRepository(conn).find_by_version(
|
||||
pg_workflow_id, current_graph_version,
|
||||
)
|
||||
]
|
||||
denied = _new_node_ref_denied(conn, previous_nodes, nodes_data, user_id)
|
||||
if denied is not None:
|
||||
return _denied(denied)
|
||||
previous_nodes = [
|
||||
serialize_node(n)
|
||||
for n in WorkflowNodesRepository(conn).find_by_version(
|
||||
pg_workflow_id, current_graph_version,
|
||||
)
|
||||
]
|
||||
# Every newly referenced node tool or source must be one the
|
||||
# caller may use, the owner included.
|
||||
denied = _new_node_ref_denied(conn, previous_nodes, nodes_data, user_id)
|
||||
if denied is not None:
|
||||
return _denied(denied)
|
||||
# A node tool/source the owner can't use runs as the editor
|
||||
# who attached it (its sponsor): only someone who owns or
|
||||
# edits it, and only once ``confirm_sponsor`` lists it.
|
||||
plan = plan_sponsors(
|
||||
conn,
|
||||
"workflow",
|
||||
workflow,
|
||||
acting,
|
||||
user_id,
|
||||
_node_refs(nodes_data),
|
||||
previous_refs=_node_refs(previous_nodes),
|
||||
confirmed=parse_confirmations(data.get("confirm_sponsor")),
|
||||
)
|
||||
refusal = sponsor_refusal(conn, "workflow", workflow, plan)
|
||||
if refusal is not None:
|
||||
body, status = refusal
|
||||
body.setdefault("error", body["message"])
|
||||
return make_response(jsonify(body), status)
|
||||
next_graph_version = current_graph_version + 1
|
||||
|
||||
_write_graph(
|
||||
@@ -695,13 +795,8 @@ class WorkflowDetail(Resource):
|
||||
"description": description,
|
||||
"current_graph_version": next_graph_version,
|
||||
}
|
||||
# A node tool/source the owner can't use runs as the editor
|
||||
# who attached it (its sponsor); record who that is.
|
||||
sponsors = sponsors_after_save(
|
||||
conn, "workflow", workflow, acting, user_id, _node_refs(nodes_data)
|
||||
)
|
||||
if sponsors != (workflow.get("resource_sponsors") or {}):
|
||||
workflow_fields["resource_sponsors"] = sponsors
|
||||
if plan.sponsors != (workflow.get("resource_sponsors") or {}):
|
||||
workflow_fields["resource_sponsors"] = plan.sponsors
|
||||
repo.update(pg_workflow_id, acting, workflow_fields)
|
||||
WorkflowNodesRepository(conn).delete_other_versions(
|
||||
pg_workflow_id, next_graph_version,
|
||||
|
||||
@@ -261,7 +261,9 @@ def chat_completions():
|
||||
internal_data["persist"] = True
|
||||
|
||||
try:
|
||||
processor = StreamProcessor(internal_data, decoded_token, trace_source="v1")
|
||||
# The token is the owner's, so tell the processor the caller is a key
|
||||
# holder: their writes on the owner's accounts need the allowlist.
|
||||
processor = StreamProcessor(internal_data, decoded_token, trace_source="v1", external_caller=True)
|
||||
flush_trace_after_request(processor)
|
||||
# Set when this request took the resume claim, so a refusal can release it.
|
||||
claimed_conversation_id = None
|
||||
|
||||
@@ -29,6 +29,7 @@ from docsgpt.api.scim import scim_bp # noqa: E402
|
||||
from docsgpt.api.user.authz import ROLE_USER, resolve_roles # noqa: E402
|
||||
from docsgpt.api.user.routes import user # noqa: E402
|
||||
from docsgpt.api.connector.routes import connector # noqa: E402
|
||||
from docsgpt.api.connector import connections as _connections # noqa: E402,F401
|
||||
from docsgpt.api.v1 import v1_bp # noqa: E402
|
||||
from docsgpt.celery_init import celery # noqa: E402
|
||||
from docsgpt.core.secret_key import resolve_jwt_secret_key # noqa: E402
|
||||
@@ -183,6 +184,27 @@ if settings.AUTH_TYPE == "simple_jwt":
|
||||
print(f"Generated Simple JWT Token: {SIMPLE_JWT_TOKEN}")
|
||||
|
||||
|
||||
def _warn_default_encryption_key() -> None:
|
||||
"""Say when stored credentials are sealed with the public default key."""
|
||||
from docsgpt.security.encryption import is_default_encryption_key
|
||||
|
||||
if not is_default_encryption_key():
|
||||
return
|
||||
if settings.AUTH_TYPE:
|
||||
logging.getLogger(__name__).warning(
|
||||
"ENCRYPTION_SECRET_KEY is the public default: connecting services is refused until you set your "
|
||||
"own value (then run `docsgpt connectors reencrypt`)."
|
||||
)
|
||||
else:
|
||||
logging.getLogger(__name__).warning(
|
||||
"ENCRYPTION_SECRET_KEY is the public default. Stored connector credentials are only as safe as "
|
||||
"that key; set your own value before exposing this install."
|
||||
)
|
||||
|
||||
|
||||
_warn_default_encryption_key()
|
||||
|
||||
|
||||
@app.route("/")
|
||||
def home():
|
||||
if request.remote_addr in ("0.0.0.0", "127.0.0.1", "localhost", "172.18.0.1"):
|
||||
@@ -209,6 +231,9 @@ def get_config():
|
||||
"hybrid_available": settings.VECTOR_STORE == "pgvector",
|
||||
"tts_available": TTSCreator.is_enabled(settings.TTS_PROVIDER),
|
||||
"stt_available": STTCreator.is_enabled(settings.STT_PROVIDER),
|
||||
# Lets a frontend built before the Connectors page run against this
|
||||
# backend, and a new frontend hide the page against an older one.
|
||||
"connectors_enabled": True,
|
||||
}
|
||||
if settings.AUTH_TYPE == "oidc":
|
||||
response["oidc"] = {
|
||||
|
||||
@@ -285,6 +285,28 @@ def _add_deploy_commands(commands) -> None:
|
||||
dev.set_defaults(func=_deploy("dev"), deploy=True)
|
||||
|
||||
|
||||
def _connectors(args: argparse.Namespace) -> int:
|
||||
"""``docsgpt connectors reencrypt``: move every credential onto the current key."""
|
||||
if getattr(args, "connectors_action", None) != "reencrypt":
|
||||
print("usage: docsgpt connectors reencrypt", file=sys.stderr)
|
||||
return 2
|
||||
from docsgpt.connectors.service import reencrypt_all
|
||||
|
||||
counts = reencrypt_all()
|
||||
print(
|
||||
f"docsgpt: re-encrypted {counts['rewritten']} connection(s), "
|
||||
f"{counts['current']} already current, {counts['failed']} unreadable",
|
||||
file=sys.stderr,
|
||||
)
|
||||
if counts["failed"]:
|
||||
print(
|
||||
"docsgpt: unreadable connections were marked 'Reconnect needed'; "
|
||||
"their owners must reconnect them.",
|
||||
file=sys.stderr,
|
||||
)
|
||||
return 1 if counts["failed"] else 0
|
||||
|
||||
|
||||
def build_parser() -> argparse.ArgumentParser:
|
||||
parser = argparse.ArgumentParser(prog="docsgpt", description="DocsGPT: private AI for agents, assistants and search.")
|
||||
parser.add_argument("--version", action="version", version=f"docsgpt {__version__}")
|
||||
@@ -315,6 +337,14 @@ def build_parser() -> argparse.ArgumentParser:
|
||||
migrate.add_argument("--no-create", dest="create_db", action="store_false", help="fail instead of creating a missing database")
|
||||
migrate.set_defaults(func=_migrate)
|
||||
|
||||
connectors = commands.add_parser("connectors", help="manage stored connector credentials")
|
||||
connector_actions = connectors.add_subparsers(dest="connectors_action", metavar="<action>")
|
||||
connector_actions.add_parser(
|
||||
"reencrypt",
|
||||
help="rewrite every stored credential with ENCRYPTION_SECRET_KEY (after a key rotation)",
|
||||
)
|
||||
connectors.set_defaults(func=_connectors)
|
||||
|
||||
for name, (module, help_text) in SCRIPTS.items():
|
||||
commands.add_parser(name, help=f"{help_text} (docsgpt.scripts.{module})", add_help=False)
|
||||
|
||||
|
||||
@@ -0,0 +1,6 @@
|
||||
"""Connectors: one place to connect accounts that feed Sources and Tools.
|
||||
|
||||
``catalog`` declares the services, ``service`` manages connections (the rows
|
||||
of ``connector_sessions``), ``permissions`` classifies tool actions as read
|
||||
or write.
|
||||
"""
|
||||
@@ -0,0 +1,68 @@
|
||||
"""Which connector a retrieved chunk came from, for "From Google Drive" citations.
|
||||
|
||||
Only the connector's key and display name are attached, never the account
|
||||
behind the connection, so shared conversations show the same attribution
|
||||
without revealing whose Drive it was.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from typing import Optional
|
||||
|
||||
from sqlalchemy import text
|
||||
|
||||
from docsgpt.storage.db.base_repository import looks_like_uuid
|
||||
from docsgpt.storage.db.session import db_readonly
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def connector_labels(source_ids: list[str]) -> dict[str, dict]:
|
||||
"""``{source_id: {"connector_key", "connector_name"}}`` for synced sources.
|
||||
|
||||
Sources that do not come from a connection are left out. Never raises:
|
||||
attribution is decoration, and a lookup failure must not break an answer.
|
||||
"""
|
||||
from docsgpt.connectors import service
|
||||
|
||||
ids = [str(i) for i in source_ids if i and looks_like_uuid(str(i))]
|
||||
if not ids:
|
||||
return {}
|
||||
try:
|
||||
with db_readonly() as conn:
|
||||
rows = conn.execute(
|
||||
text(
|
||||
"SELECT s.id AS source_id, cs.* FROM sources s "
|
||||
"JOIN connector_sessions cs ON cs.id = s.connection_id "
|
||||
"WHERE s.id = ANY(CAST(:ids AS uuid[]))"
|
||||
),
|
||||
{"ids": ids},
|
||||
).fetchall()
|
||||
except Exception:
|
||||
logger.warning("connector attribution lookup failed", exc_info=True)
|
||||
return {}
|
||||
labels = {}
|
||||
for row in rows:
|
||||
data = dict(row._mapping)
|
||||
public = service.serialize_connection({**data, "id": data["id"]})
|
||||
labels[str(data["source_id"])] = {
|
||||
"connector_key": public["connector_key"],
|
||||
"connector_name": public["name"],
|
||||
}
|
||||
return labels
|
||||
|
||||
|
||||
class ConnectorLabelCache:
|
||||
"""Per-search memo so each source is looked up once."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._labels: dict[str, dict] = {}
|
||||
|
||||
def for_source(self, source_id: Optional[str]) -> dict:
|
||||
if not source_id:
|
||||
return {}
|
||||
key = str(source_id)
|
||||
if key not in self._labels:
|
||||
self._labels[key] = connector_labels([key]).get(key, {})
|
||||
return self._labels[key]
|
||||
@@ -0,0 +1,522 @@
|
||||
"""The connector catalog: every service DocsGPT can connect to.
|
||||
|
||||
A connector is something the user connects once (an OAuth sign-in, an API key
|
||||
or an MCP server). Each connection it produces can feed Sources (content
|
||||
synced into DocsGPT) and Tools (actions an agent can take). The definitions
|
||||
here are declarative: they say how a connector signs in, which server
|
||||
settings it needs, what it can sync and which tools it creates, so the API
|
||||
and the frontend never hard-code a list of services.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
from typing import Any, Iterable, Optional
|
||||
from urllib.parse import urlparse
|
||||
|
||||
import yaml
|
||||
|
||||
from docsgpt.core.settings import settings
|
||||
|
||||
CATEGORIES = ("files", "knowledge", "projects", "dev", "business", "messaging", "database", "search", "custom")
|
||||
AUTH_KINDS = ("oauth", "mcp_oauth", "api_key", "none", "mcp")
|
||||
PUBLISHERS = ("built_in", "preset", "custom")
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class CredentialField:
|
||||
"""One field the user fills in to connect an ``api_key`` connector.
|
||||
|
||||
Attributes:
|
||||
key: Name the credential is stored and passed to the tool or loader
|
||||
under, e.g. ``token`` or ``aws_access_key_id``.
|
||||
label: English label; the frontend shows the translated
|
||||
``settings.connectors.fields.<key>`` when it has one.
|
||||
secret: Masked in the form and never returned by the API.
|
||||
required: Whether connecting fails without it.
|
||||
parameter: A tool parameter this field sets. When the connection has
|
||||
a value for it, every call through the connection uses that
|
||||
value and the model is not asked for the parameter (Telegram's
|
||||
default chat).
|
||||
hint: English help shown under the field; the frontend shows the
|
||||
translated ``settings.connectors.fieldHints.<connector>_<key>``.
|
||||
"""
|
||||
|
||||
key: str
|
||||
label: str
|
||||
secret: bool = True
|
||||
required: bool = True
|
||||
parameter: Optional[str] = None
|
||||
hint: Optional[str] = None
|
||||
|
||||
def to_dict(self) -> dict:
|
||||
return {
|
||||
"key": self.key,
|
||||
"label": self.label,
|
||||
"secret": self.secret,
|
||||
"required": self.required,
|
||||
"parameter": self.parameter,
|
||||
"hint": self.hint,
|
||||
}
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ConnectorDefinition:
|
||||
"""How one connector signs in and what a connection to it provides.
|
||||
|
||||
Attributes:
|
||||
key: Catalog key, stored on each connection as ``connector_key``.
|
||||
name: Display name.
|
||||
description: One line shown on the catalog card.
|
||||
icon: Frontend asset key (``assets/connectors/<icon>.svg``).
|
||||
category: One of :data:`CATEGORIES`.
|
||||
auth_kind: ``oauth`` (server-side OAuth app), ``mcp_oauth`` (MCP
|
||||
OAuth with dynamic client registration), ``api_key`` (fields the
|
||||
user pastes), ``none`` or ``mcp`` (a custom MCP server whose auth
|
||||
the user picks).
|
||||
capabilities: Any of ``sync``, ``read``, ``write``.
|
||||
credential_fields: Fields asked for by ``api_key`` connectors.
|
||||
required_settings: Server settings that must be set for the
|
||||
connector to be usable, e.g. ``GOOGLE_CLIENT_ID``.
|
||||
sync_ingestor: Ingest loader a synced source uses (``google_drive``,
|
||||
``s3``), or None when the connector cannot sync.
|
||||
default_sync_frequency: Sync frequency preselected in the wizard.
|
||||
setup_fields: Per-source fields asked when choosing what to sync
|
||||
(an S3 bucket, Reddit search queries).
|
||||
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, and for a built-in connector whose
|
||||
tool is its service's own MCP server (GitHub).
|
||||
mcp_write_url: The same server's endpoint that also offers write
|
||||
actions, which a connection opts into and an admin can forbid
|
||||
(GitHub's full server beside its read-only one). ``mcp_url``
|
||||
stays the default.
|
||||
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
|
||||
name: str
|
||||
description: str
|
||||
icon: str
|
||||
category: str
|
||||
auth_kind: str
|
||||
capabilities: tuple[str, ...] = ()
|
||||
credential_fields: tuple[CredentialField, ...] = ()
|
||||
required_settings: tuple[str, ...] = ()
|
||||
sync_ingestor: Optional[str] = None
|
||||
default_sync_frequency: str = "weekly"
|
||||
setup_fields: tuple[CredentialField, ...] = ()
|
||||
tool_templates: tuple[str, ...] = ()
|
||||
setup: dict = field(default_factory=lambda: {"tools": "auto", "sync": "ask"})
|
||||
mcp_url: Optional[str] = None
|
||||
mcp_write_url: Optional[str] = None
|
||||
publisher: str = "built_in"
|
||||
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]:
|
||||
"""Server settings this connector needs that are not set."""
|
||||
return [name for name in self.required_settings if not getattr(settings, name, None)]
|
||||
|
||||
@property
|
||||
def configured(self) -> bool:
|
||||
"""Whether every required server setting is set."""
|
||||
return not self.missing_settings
|
||||
|
||||
@property
|
||||
def mcp_base_url(self) -> Optional[str]:
|
||||
"""``scheme://host`` of the preset's MCP endpoint (the MCP session key)."""
|
||||
return base_url(self.mcp_url) if self.mcp_url else None
|
||||
|
||||
def to_dict(self) -> dict:
|
||||
"""Serialise the parts the frontend needs (never server secrets)."""
|
||||
return {
|
||||
"key": self.key,
|
||||
"name": self.name,
|
||||
"description": self.description,
|
||||
"icon": self.icon,
|
||||
"category": self.category,
|
||||
"auth_kind": self.auth_kind,
|
||||
"capabilities": list(self.capabilities),
|
||||
"credential_fields": [f.to_dict() for f in self.credential_fields],
|
||||
"setup_fields": [f.to_dict() for f in self.setup_fields],
|
||||
"sync_ingestor": self.sync_ingestor,
|
||||
"default_sync_frequency": self.default_sync_frequency,
|
||||
"tool_templates": list(self.tool_templates),
|
||||
"setup": dict(self.setup),
|
||||
"mcp_url": self.mcp_url,
|
||||
"writes_opt_in": bool(self.mcp_write_url),
|
||||
"publisher": self.publisher,
|
||||
"docs_url": self.docs_url,
|
||||
"oauth_scopes": list(self.oauth_scopes),
|
||||
"part_of": self.part_of,
|
||||
"sign_in_methods": self.sign_in_methods,
|
||||
}
|
||||
|
||||
|
||||
def base_url(url: Optional[str]) -> str:
|
||||
"""``scheme://netloc`` of ``url``, the key MCP sessions are stored under."""
|
||||
parsed = urlparse(url or "")
|
||||
if not parsed.scheme or not parsed.netloc:
|
||||
return ""
|
||||
return f"{parsed.scheme}://{parsed.netloc}"
|
||||
|
||||
|
||||
_DOCS = "https://docs.docsgpt.cloud/Guides/Connectors"
|
||||
GITHUB_MCP_URL = "https://api.githubcopilot.com/mcp/readonly"
|
||||
GITHUB_MCP_WRITE_URL = "https://api.githubcopilot.com/mcp/"
|
||||
|
||||
# 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(
|
||||
key="google_drive",
|
||||
name="Google Drive",
|
||||
description="Sync Docs, Sheets and PDFs into Knowledge.",
|
||||
icon="drive",
|
||||
category="files",
|
||||
auth_kind="oauth",
|
||||
capabilities=("sync",),
|
||||
required_settings=("GOOGLE_CLIENT_ID", "GOOGLE_CLIENT_SECRET"),
|
||||
sync_ingestor="google_drive",
|
||||
setup={"tools": "off", "sync": "ask"},
|
||||
docs_url=f"{_DOCS}#google-drive",
|
||||
),
|
||||
ConnectorDefinition(
|
||||
key="share_point",
|
||||
name="SharePoint",
|
||||
description="Sync files from SharePoint sites and OneDrive into Knowledge.",
|
||||
icon="sharepoint",
|
||||
category="files",
|
||||
auth_kind="oauth",
|
||||
capabilities=("sync",),
|
||||
required_settings=("MICROSOFT_CLIENT_ID", "MICROSOFT_CLIENT_SECRET"),
|
||||
sync_ingestor="share_point",
|
||||
setup={"tools": "off", "sync": "ask"},
|
||||
docs_url=f"{_DOCS}#sharepoint-and-onedrive",
|
||||
),
|
||||
ConnectorDefinition(
|
||||
key="confluence",
|
||||
name="Confluence",
|
||||
description="Sync Confluence spaces and pages into Knowledge.",
|
||||
icon="confluence",
|
||||
category="knowledge",
|
||||
auth_kind="oauth",
|
||||
capabilities=("sync",),
|
||||
required_settings=("CONFLUENCE_CLIENT_ID", "CONFLUENCE_CLIENT_SECRET"),
|
||||
sync_ingestor="confluence",
|
||||
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 unless the connection opts into
|
||||
# changes (issues, comments, pull requests) and an admin allows them.
|
||||
tool_templates=("mcp_tool",),
|
||||
mcp_url=GITHUB_MCP_URL,
|
||||
mcp_write_url=GITHUB_MCP_WRITE_URL,
|
||||
setup={"tools": "ask", "sync": "ask"},
|
||||
docs_url=f"{_DOCS}#github",
|
||||
),
|
||||
ConnectorDefinition(
|
||||
key="s3",
|
||||
name="Amazon S3",
|
||||
description="Sync documents from an S3 bucket into Knowledge.",
|
||||
icon="s3",
|
||||
category="files",
|
||||
auth_kind="api_key",
|
||||
capabilities=("sync",),
|
||||
credential_fields=(
|
||||
CredentialField("aws_access_key_id", "Access key ID", secret=False),
|
||||
CredentialField("aws_secret_access_key", "Secret access key"),
|
||||
CredentialField("region", "Region", secret=False, required=False),
|
||||
CredentialField("endpoint_url", "Custom endpoint URL", secret=False, required=False),
|
||||
),
|
||||
sync_ingestor="s3",
|
||||
setup_fields=(
|
||||
CredentialField("bucket", "Bucket", secret=False),
|
||||
CredentialField("prefix", "Path prefix", secret=False, required=False),
|
||||
),
|
||||
setup={"tools": "off", "sync": "ask"},
|
||||
docs_url=f"{_DOCS}#amazon-s3",
|
||||
),
|
||||
ConnectorDefinition(
|
||||
key="reddit",
|
||||
name="Reddit",
|
||||
description="Sync Reddit posts that match your searches into Knowledge.",
|
||||
icon="reddit",
|
||||
category="search",
|
||||
auth_kind="api_key",
|
||||
capabilities=("sync",),
|
||||
credential_fields=(
|
||||
CredentialField("client_id", "Client ID", secret=False),
|
||||
CredentialField("client_secret", "Client secret"),
|
||||
CredentialField("user_agent", "User agent", secret=False),
|
||||
),
|
||||
sync_ingestor="reddit",
|
||||
setup_fields=(
|
||||
CredentialField("search_queries", "Search queries", secret=False),
|
||||
CredentialField("number_posts", "Number of posts", secret=False),
|
||||
),
|
||||
setup={"tools": "off", "sync": "ask"},
|
||||
),
|
||||
ConnectorDefinition(
|
||||
key="brave",
|
||||
name="Brave Search",
|
||||
description="Search the web and images with the Brave Search API.",
|
||||
icon="tool_brave",
|
||||
category="search",
|
||||
auth_kind="api_key",
|
||||
capabilities=("read",),
|
||||
credential_fields=(CredentialField("token", "API key"),),
|
||||
tool_templates=("brave",),
|
||||
setup={"tools": "auto", "sync": "off"},
|
||||
),
|
||||
ConnectorDefinition(
|
||||
key="telegram",
|
||||
name="Telegram",
|
||||
description="Send messages and images to a Telegram chat.",
|
||||
icon="tool_telegram",
|
||||
category="messaging",
|
||||
auth_kind="api_key",
|
||||
capabilities=("write",),
|
||||
credential_fields=(
|
||||
CredentialField("token", "Bot token"),
|
||||
CredentialField(
|
||||
"chat_id",
|
||||
"Default chat ID",
|
||||
secret=False,
|
||||
required=False,
|
||||
parameter="chat_id",
|
||||
hint=(
|
||||
"Optional. Messages go to this chat, and the AI cannot pick another. Add the bot to the "
|
||||
"chat, send it a message, then find the chat's id in "
|
||||
"https://api.telegram.org/bot<token>/getUpdates."
|
||||
),
|
||||
),
|
||||
),
|
||||
tool_templates=("telegram",),
|
||||
setup={"tools": "auto", "sync": "off"},
|
||||
),
|
||||
ConnectorDefinition(
|
||||
key="ntfy",
|
||||
name="ntfy",
|
||||
description="Send push notifications through an ntfy server.",
|
||||
icon="tool_ntfy",
|
||||
category="messaging",
|
||||
auth_kind="api_key",
|
||||
capabilities=("write",),
|
||||
credential_fields=(CredentialField("token", "Access token"),),
|
||||
tool_templates=("ntfy",),
|
||||
setup={"tools": "auto", "sync": "off"},
|
||||
),
|
||||
ConnectorDefinition(
|
||||
key="postgres",
|
||||
name="PostgreSQL",
|
||||
description="Read the schema and run SQL against a Postgres database.",
|
||||
icon="tool_postgres",
|
||||
category="database",
|
||||
auth_kind="api_key",
|
||||
capabilities=("read", "write"),
|
||||
credential_fields=(CredentialField("token", "Connection string"),),
|
||||
tool_templates=("postgres",),
|
||||
setup={"tools": "auto", "sync": "off"},
|
||||
),
|
||||
ConnectorDefinition(
|
||||
key="custom_mcp",
|
||||
name="MCP server",
|
||||
description="Connect any remote Model Context Protocol server.",
|
||||
icon="tool_mcp_tool",
|
||||
category="custom",
|
||||
auth_kind="mcp",
|
||||
capabilities=("read", "write"),
|
||||
tool_templates=("mcp_tool",),
|
||||
setup={"tools": "auto", "sync": "off"},
|
||||
publisher="custom",
|
||||
),
|
||||
ConnectorDefinition(
|
||||
key="custom_openapi",
|
||||
name="OpenAPI / REST",
|
||||
description="Import an OpenAPI spec and call its endpoints as tools.",
|
||||
icon="tool_api_tool",
|
||||
category="custom",
|
||||
auth_kind="none",
|
||||
capabilities=("read", "write"),
|
||||
tool_templates=("api_tool",),
|
||||
setup={"tools": "ask", "sync": "off"},
|
||||
publisher="custom",
|
||||
),
|
||||
)
|
||||
|
||||
_PRESETS_FILE = Path(__file__).parent / "presets" / "mcp.yaml"
|
||||
|
||||
|
||||
def _load_presets(path: Optional[Path] = None) -> tuple[ConnectorDefinition, ...]:
|
||||
"""Read the curated MCP server presets shipped with the repository."""
|
||||
path = path or _PRESETS_FILE
|
||||
if not path.exists():
|
||||
return ()
|
||||
with path.open(encoding="utf-8") as fh:
|
||||
entries = yaml.safe_load(fh) or []
|
||||
presets = []
|
||||
for entry in entries:
|
||||
presets.append(
|
||||
ConnectorDefinition(
|
||||
key=entry["key"],
|
||||
name=entry["name"],
|
||||
description=entry["description"],
|
||||
icon=entry.get("icon") or "tool_mcp_tool",
|
||||
category=entry.get("category", "knowledge"),
|
||||
auth_kind=entry.get("auth_kind", "mcp_oauth"),
|
||||
capabilities=tuple(entry.get("capabilities") or ("read", "write")),
|
||||
tool_templates=("mcp_tool",),
|
||||
# Linear also syncs into Knowledge, read through its MCP server.
|
||||
sync_ingestor=entry.get("sync_ingestor"),
|
||||
setup={"tools": "auto", "sync": "ask" if entry.get("sync_ingestor") else "off"},
|
||||
mcp_url=entry["mcp_url"],
|
||||
publisher="preset",
|
||||
docs_url=entry.get("docs_url"),
|
||||
oauth_scopes=tuple(entry.get("oauth_scopes") or ()),
|
||||
part_of=entry.get("part_of"),
|
||||
)
|
||||
)
|
||||
return tuple(presets)
|
||||
|
||||
|
||||
_REGISTRY: dict[str, ConnectorDefinition] = {}
|
||||
|
||||
|
||||
def _registry() -> dict[str, ConnectorDefinition]:
|
||||
if not _REGISTRY:
|
||||
for definition in (*_BUILT_IN, *_load_presets()):
|
||||
_REGISTRY[definition.key] = definition
|
||||
return _REGISTRY
|
||||
|
||||
|
||||
def all_definitions() -> list[ConnectorDefinition]:
|
||||
"""Every catalog entry, built-ins first, then presets."""
|
||||
return list(_registry().values())
|
||||
|
||||
|
||||
def get_definition(key: Optional[str]) -> Optional[ConnectorDefinition]:
|
||||
"""The definition for ``key``, or None when it is not in the catalog."""
|
||||
if not key:
|
||||
return None
|
||||
return _registry().get(key)
|
||||
|
||||
|
||||
def preset_for_url(url: Optional[str]) -> Optional[ConnectorDefinition]:
|
||||
"""The MCP preset whose server shares ``url``'s base URL, if any."""
|
||||
target = base_url(url)
|
||||
if not target:
|
||||
return None
|
||||
for definition in _registry().values():
|
||||
if definition.publisher == "preset" and definition.mcp_base_url == target:
|
||||
return definition
|
||||
return None
|
||||
|
||||
|
||||
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
|
||||
return None
|
||||
|
||||
|
||||
def parameter_fields(key: Optional[str]) -> tuple[CredentialField, ...]:
|
||||
"""The credential fields of connector ``key`` that set a tool parameter."""
|
||||
definition = get_definition(key)
|
||||
if definition is None:
|
||||
return ()
|
||||
return tuple(f for f in definition.credential_fields if f.parameter)
|
||||
|
||||
|
||||
def connector_key_for_row(row: dict) -> Optional[str]:
|
||||
"""Catalog key for a ``connector_sessions`` row.
|
||||
|
||||
Rows written before ``0040_connections`` carry no key; they are named from
|
||||
``provider``. A custom MCP row whose server matches a preset is reported
|
||||
as that preset.
|
||||
"""
|
||||
key = row.get("connector_key")
|
||||
provider = row.get("provider") or ""
|
||||
if not key or key.startswith("mcp:"):
|
||||
# ``mcp:<server>`` is a provider value, never a catalog key.
|
||||
key = None
|
||||
if provider.startswith("mcp:"):
|
||||
key = "custom_mcp"
|
||||
elif get_definition(provider):
|
||||
key = provider
|
||||
if key == "custom_mcp":
|
||||
preset = preset_for_url(row.get("server_url") or provider[4:])
|
||||
if preset:
|
||||
return preset.key
|
||||
return key
|
||||
|
||||
|
||||
def tool_connector_keys() -> set[str]:
|
||||
"""``user_tools`` names that belong to a built-in service connector."""
|
||||
return {
|
||||
name
|
||||
for definition in _BUILT_IN
|
||||
if definition.publisher == "built_in"
|
||||
for name in definition.tool_templates
|
||||
if name not in _GENERIC_TOOL_TEMPLATES
|
||||
}
|
||||
|
||||
|
||||
def iter_by_category(definitions: Iterable[ConnectorDefinition]) -> dict[str, list[ConnectorDefinition]]:
|
||||
"""Group definitions by category, keeping :data:`CATEGORIES` order."""
|
||||
grouped: dict[str, list[ConnectorDefinition]] = {c: [] for c in CATEGORIES}
|
||||
for definition in definitions:
|
||||
grouped.setdefault(definition.category, []).append(definition)
|
||||
return grouped
|
||||
|
||||
|
||||
def reset_registry_for_tests() -> None:
|
||||
"""Drop the cached registry so a test can load a different presets file."""
|
||||
_REGISTRY.clear()
|
||||
|
||||
|
||||
def to_public(definition: ConnectorDefinition, **extra: Any) -> dict:
|
||||
"""Definition dict plus request-specific fields (availability, counts)."""
|
||||
return {**definition.to_dict(), **extra}
|
||||
@@ -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]
|
||||
@@ -0,0 +1,299 @@
|
||||
"""Linear through its MCP server: what a sync can pick, and reading issues and documents.
|
||||
|
||||
Linear's hosted MCP server (``mcp.linear.app``) is its own OAuth issuer, and
|
||||
the tokens it grants are for that server. Knowledge sync therefore reads
|
||||
Linear through the same MCP tools the agents use (``list_issues``,
|
||||
``get_issue``, ``list_comments``, ``list_documents``...), signed in with the
|
||||
Linear connection, so one sign-in powers both and nobody registers an OAuth
|
||||
app. Linear publishes no output schema for these tools, so records are read
|
||||
defensively: a value may be a string or an object with a ``name``.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import re
|
||||
from typing import Any, AsyncIterator, Optional
|
||||
|
||||
from docsgpt.connectors import catalog
|
||||
|
||||
LINEAR_CONNECTOR = "mcp:linear"
|
||||
LINEAR_MCP_URL = "https://mcp.linear.app/mcp"
|
||||
# Enough for a team's working set; a sync is a full read, so it stays bounded.
|
||||
MAX_ISSUES = 500
|
||||
MAX_DOCUMENTS = 100
|
||||
MAX_COMMENTS = 100
|
||||
# Teams and projects offered in the picker.
|
||||
MAX_PICKER_ITEMS = 250
|
||||
PAGE_SIZE = 100
|
||||
_MAX_PAGES = 50
|
||||
|
||||
_IDENTIFIER = re.compile(r"^[A-Za-z][A-Za-z0-9_]*-\d+$")
|
||||
_TRUNCATED = re.compile(r"\(truncated\b", re.IGNORECASE)
|
||||
_LIST_KEYS = ("nodes", "items", "results", "data")
|
||||
_PRIORITIES = {0: "No priority", 1: "Urgent", 2: "High", 3: "Medium", 4: "Low"}
|
||||
|
||||
|
||||
class LinearSyncError(ValueError):
|
||||
"""Linear's MCP server cannot do what the sync needs (a tool or filter is gone)."""
|
||||
|
||||
|
||||
def mcp_url() -> str:
|
||||
"""The Linear preset's MCP endpoint."""
|
||||
definition = catalog.get_definition(LINEAR_CONNECTOR)
|
||||
return (definition.mcp_url if definition and definition.mcp_url else None) or LINEAR_MCP_URL
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# What a source syncs
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _truthy(value: Any, default: bool) -> bool:
|
||||
if value is None:
|
||||
return default
|
||||
if isinstance(value, str):
|
||||
return value.strip().lower() not in ("", "0", "false", "no", "off")
|
||||
return bool(value)
|
||||
|
||||
|
||||
def _picked(values: Any, fields: tuple[str, ...]) -> list[dict]:
|
||||
picked: dict[str, dict] = {}
|
||||
for value in values if isinstance(values, list) else []:
|
||||
record = value if isinstance(value, dict) else {"id": value}
|
||||
item_id = str(record.get("id") or "").strip()
|
||||
if not item_id:
|
||||
continue
|
||||
current = picked.setdefault(item_id, {"id": item_id, **{f: "" for f in fields}})
|
||||
for field in fields:
|
||||
current[field] = current[field] or str(record.get(field) or "").strip()
|
||||
return list(picked.values())
|
||||
|
||||
|
||||
def normalize_selection(items: Any) -> dict:
|
||||
"""What a Linear source syncs, as it is stored in ``sources.remote_data``.
|
||||
|
||||
Args:
|
||||
items: The wizard's choice (or a stored ``remote_data``, possibly as
|
||||
JSON): ``teams`` and ``projects`` as ids or ``{id, key, name}``
|
||||
records, ``include_comments`` (on unless turned off) and
|
||||
``include_documents`` (the picked projects' documents).
|
||||
|
||||
Returns:
|
||||
``{teams: [{id, key, name}], projects: [{id, name}],
|
||||
include_comments, include_documents}``.
|
||||
|
||||
Raises:
|
||||
ValueError: Nothing to sync was picked.
|
||||
"""
|
||||
if isinstance(items, str):
|
||||
try:
|
||||
items = json.loads(items)
|
||||
except ValueError:
|
||||
items = {}
|
||||
items = items if isinstance(items, dict) else {}
|
||||
selection = {
|
||||
"teams": _picked(items.get("teams"), ("key", "name")),
|
||||
"projects": _picked(items.get("projects"), ("name",)),
|
||||
"include_comments": _truthy(items.get("include_comments"), True),
|
||||
"include_documents": _truthy(items.get("include_documents"), False),
|
||||
}
|
||||
if not selection["teams"] and not selection["projects"]:
|
||||
raise ValueError("Pick at least one Linear team or project")
|
||||
return selection
|
||||
|
||||
|
||||
def selection_name(selection: dict) -> str:
|
||||
"""``Linear · Engineering, Acme``: a source's default name, from what it syncs."""
|
||||
names = [item.get("name") for item in (*selection["teams"], *selection["projects"]) if item.get("name")]
|
||||
return f"Linear · {', '.join(names)}" if names else "Linear"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Reading tool answers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def items_of(payload: Any, key: str) -> list[dict]:
|
||||
"""The records in a list tool's answer: ``{key: [...]}``, a GraphQL ``nodes`` list or a bare list."""
|
||||
if isinstance(payload, list):
|
||||
found = payload
|
||||
elif isinstance(payload, dict):
|
||||
found = payload.get(key)
|
||||
if isinstance(found, dict):
|
||||
found = found.get("nodes")
|
||||
if not isinstance(found, list):
|
||||
found = next((payload[k] for k in _LIST_KEYS if isinstance(payload.get(k), list)), None)
|
||||
if found is None:
|
||||
lists = [value for value in payload.values() if isinstance(value, list)]
|
||||
found = lists[0] if len(lists) == 1 else []
|
||||
else:
|
||||
found = []
|
||||
return [item for item in found if isinstance(item, dict)]
|
||||
|
||||
|
||||
def next_cursor(payload: Any) -> Optional[str]:
|
||||
"""The cursor of the next page, or None on the last one."""
|
||||
if not isinstance(payload, dict):
|
||||
return None
|
||||
info = payload.get("pageInfo") if isinstance(payload.get("pageInfo"), dict) else payload
|
||||
if info.get("hasNextPage") is False:
|
||||
return None
|
||||
cursor = info.get("cursor") or info.get("nextCursor") or info.get("endCursor")
|
||||
return str(cursor) if cursor else None
|
||||
|
||||
|
||||
def unwrap(payload: Any, key: str) -> dict:
|
||||
"""A single record from a get tool's answer, bare or as ``{key: {...}}``."""
|
||||
if isinstance(payload, dict) and isinstance(payload.get(key), dict):
|
||||
return payload[key]
|
||||
return payload if isinstance(payload, dict) else {}
|
||||
|
||||
|
||||
def name_of(value: Any) -> str:
|
||||
"""A person's, state's or label's name, whether Linear sent a string or an object."""
|
||||
if isinstance(value, dict):
|
||||
value = value.get("name") or value.get("displayName") or value.get("title") or value.get("key")
|
||||
return str(value).strip() if value not in (None, "") else ""
|
||||
|
||||
|
||||
def names_of(values: Any) -> list[str]:
|
||||
"""Label names from a list of strings or objects (or a ``{nodes: [...]}``)."""
|
||||
if isinstance(values, dict):
|
||||
values = values.get("nodes")
|
||||
return [name for name in (name_of(value) for value in values or []) if name] if isinstance(values, list) else []
|
||||
|
||||
|
||||
def priority_of(value: Any) -> str:
|
||||
"""``High``: a priority given as a label, an object or Linear's 0-4 number."""
|
||||
if isinstance(value, dict):
|
||||
return str(value.get("name") or _PRIORITIES.get(value.get("value"), "")).strip()
|
||||
if isinstance(value, bool):
|
||||
return ""
|
||||
if isinstance(value, int):
|
||||
return _PRIORITIES.get(value, "") if value else ""
|
||||
return str(value or "").strip()
|
||||
|
||||
|
||||
def issue_identifier(issue: dict) -> str:
|
||||
"""``ENG-123``: the identifier people use, which list tools may send as the ``id``."""
|
||||
identifier = issue.get("identifier")
|
||||
if identifier:
|
||||
return str(identifier)
|
||||
item_id = str(issue.get("id") or "")
|
||||
return item_id if _IDENTIFIER.match(item_id) else ""
|
||||
|
||||
|
||||
def is_truncated(text: Any) -> bool:
|
||||
"""Whether Linear clipped a description in a list (it marks it ``(truncated…)``)."""
|
||||
return isinstance(text, str) and bool(_TRUNCATED.search(text))
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Calling the tools
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
async def _schema(session: Any, tool: str) -> dict:
|
||||
schema = await session.input_schema(tool)
|
||||
if schema is None:
|
||||
raise LinearSyncError(f"Linear's MCP server no longer offers {tool}")
|
||||
return schema
|
||||
|
||||
|
||||
def argument(schema: dict, *names: str) -> Optional[str]:
|
||||
"""The first of ``names`` the tool takes; the first name when the schema lists no properties."""
|
||||
properties = schema.get("properties") if isinstance(schema, dict) else None
|
||||
if not isinstance(properties, dict) or not properties:
|
||||
return names[0]
|
||||
return next((name for name in names if name in properties), None)
|
||||
|
||||
|
||||
def _page_size(schema: dict, name: str) -> int:
|
||||
spec = (schema.get("properties") or {}).get(name) if isinstance(schema, dict) else None
|
||||
maximum = spec.get("maximum") if isinstance(spec, dict) else None
|
||||
return min(PAGE_SIZE, int(maximum)) if isinstance(maximum, (int, float)) and maximum > 0 else PAGE_SIZE
|
||||
|
||||
|
||||
async def paged(
|
||||
session: Any,
|
||||
tool: str,
|
||||
key: str,
|
||||
filters: dict,
|
||||
*,
|
||||
limit: int,
|
||||
extra: Optional[dict] = None,
|
||||
) -> AsyncIterator[dict]:
|
||||
"""Every record a list tool returns for ``filters``, page by page, up to ``limit``.
|
||||
|
||||
Args:
|
||||
session: The open MCP session.
|
||||
tool: The list tool, e.g. ``list_issues``.
|
||||
key: Where its answer holds the records, e.g. ``issues``.
|
||||
filters: Filters the sync depends on, as ``{(name, alias...): value}``.
|
||||
A filter the tool does not take is an error, never dropped: a
|
||||
team's sync must not read the whole workspace.
|
||||
limit: Most records to yield.
|
||||
extra: Optional arguments, sent only when the tool takes them.
|
||||
|
||||
Raises:
|
||||
LinearSyncError: The tool is gone or cannot apply a filter.
|
||||
"""
|
||||
schema = await _schema(session, tool)
|
||||
arguments: dict = {}
|
||||
for names, value in filters.items():
|
||||
name = argument(schema, *names)
|
||||
if name is None:
|
||||
raise LinearSyncError(f"Linear's {tool} can no longer filter by {names[0]}")
|
||||
arguments[name] = value
|
||||
size_name = argument(schema, "limit", "first")
|
||||
if size_name:
|
||||
arguments[size_name] = _page_size(schema, size_name)
|
||||
properties = schema.get("properties") if isinstance(schema, dict) else None
|
||||
for name, value in (extra or {}).items():
|
||||
if isinstance(properties, dict) and name in properties:
|
||||
arguments[name] = value
|
||||
cursor_name = argument(schema, "cursor", "after")
|
||||
cursor: Optional[str] = None
|
||||
seen: set[str] = set()
|
||||
yielded = 0
|
||||
for _ in range(_MAX_PAGES):
|
||||
payload = await session.call(tool, {**arguments, **({cursor_name: cursor} if cursor else {})})
|
||||
for record in items_of(payload, key):
|
||||
yield record
|
||||
yielded += 1
|
||||
if yielded >= limit:
|
||||
return
|
||||
cursor = next_cursor(payload)
|
||||
if not cursor or not cursor_name or cursor in seen:
|
||||
return
|
||||
seen.add(cursor)
|
||||
|
||||
|
||||
async def list_workspace(session: Any) -> dict:
|
||||
"""The teams and projects a Linear account can see, for the sync picker.
|
||||
|
||||
Returns:
|
||||
``{teams: [{id, key, name}], projects: [{id, name, state, teams}]}``,
|
||||
each sorted by name.
|
||||
"""
|
||||
teams = []
|
||||
async for team in paged(session, "list_teams", "teams", {}, limit=MAX_PICKER_ITEMS):
|
||||
if team.get("id"):
|
||||
teams.append({"id": str(team["id"]), "key": str(team.get("key") or ""),
|
||||
"name": name_of(team.get("name")) or str(team.get("key") or team["id"])})
|
||||
projects = []
|
||||
async for project in paged(session, "list_projects", "projects", {}, limit=MAX_PICKER_ITEMS,
|
||||
extra={"includeArchived": False}):
|
||||
if project.get("id"):
|
||||
project_teams = project.get("teams") or ([project["team"]] if project.get("team") else [])
|
||||
projects.append({
|
||||
"id": str(project["id"]),
|
||||
"name": name_of(project.get("name")) or str(project["id"]),
|
||||
"state": name_of(project.get("state") or project.get("status")),
|
||||
"teams": names_of(project_teams),
|
||||
})
|
||||
return {
|
||||
"teams": sorted(teams, key=lambda t: t["name"].lower()),
|
||||
"projects": sorted(projects, key=lambda p: p["name"].lower()),
|
||||
}
|
||||
@@ -0,0 +1,382 @@
|
||||
"""MCP helpers for connections: re-scanning a server's actions, switching GitHub's
|
||||
writes, and calling a server's tools.
|
||||
|
||||
Agents call an MCP server's tools through ``MCPTool``. Server-side work that
|
||||
reads a service through its MCP server (syncing Linear into Knowledge) opens
|
||||
one session with :func:`run_connection_session` and makes all its calls in
|
||||
it, signed in with the connection's own OAuth tokens.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import concurrent.futures
|
||||
import json
|
||||
from typing import Any, Awaitable, Callable, Optional, TypeVar
|
||||
|
||||
from docsgpt.agents.tool_pins import carry_pins
|
||||
from docsgpt.connectors import catalog, service
|
||||
from docsgpt.connectors.catalog import ConnectorDefinition
|
||||
from docsgpt.connectors.permissions import ACCESS_READ, ACCESS_WRITE, apply_default_permissions
|
||||
from docsgpt.storage.db.repositories.connector_sessions import ConnectorSessionsRepository
|
||||
from docsgpt.storage.db.repositories.user_tools import UserToolsRepository
|
||||
from docsgpt.storage.db.session import db_readonly, db_session
|
||||
|
||||
T = TypeVar("T")
|
||||
|
||||
|
||||
class NoMcpTool(LookupError):
|
||||
"""The connection has no MCP tool yet (it was set up without tools)."""
|
||||
|
||||
|
||||
def _builtin_definition(connection: dict) -> Optional[ConnectorDefinition]:
|
||||
"""The connection's built-in connector when its tool is that service's MCP server."""
|
||||
definition = catalog.get_definition(catalog.connector_key_for_row(connection))
|
||||
return definition if service.builtin_mcp_config(definition) is not None else None
|
||||
|
||||
|
||||
def _write_endpoint_access(action: dict) -> dict:
|
||||
"""Stamp an action of a write endpoint: a read only when the server marks it read-only.
|
||||
|
||||
GitHub marks its read tools, so a name is no evidence there: from the name
|
||||
alone ``mark_all_notifications_read`` would pass for a read.
|
||||
"""
|
||||
annotations = action.get("annotations") if isinstance(action.get("annotations"), dict) else {}
|
||||
return {**action, "access": ACCESS_READ if annotations.get("readOnlyHint") is True else ACCESS_WRITE}
|
||||
|
||||
|
||||
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", "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"])
|
||||
config["query_mode"] = True
|
||||
mcp_tool = MCPTool(config, user_id)
|
||||
mcp_tool.discover_tools()
|
||||
return mcp_tool.get_actions_metadata()
|
||||
|
||||
|
||||
def _classified(connection: dict, config: dict, actions: list[dict]) -> list[dict]:
|
||||
"""``actions`` as read from ``config``'s server, stamped strictly when it is a write endpoint."""
|
||||
definition = _builtin_definition(connection)
|
||||
if definition is None or not definition.mcp_write_url or config.get("server_url") != definition.mcp_write_url:
|
||||
return actions
|
||||
return [_write_endpoint_access(action) for action in actions]
|
||||
|
||||
|
||||
def discover_builtin_actions(user_id: str, connection: dict, writes: bool = False) -> list[dict]:
|
||||
"""The actions of a built-in connector's MCP server (GitHub's), read with the connection.
|
||||
|
||||
Args:
|
||||
user_id: The connection's owner.
|
||||
connection: The connection row.
|
||||
writes: Read the endpoint that also offers write actions; the caller
|
||||
checked that an admin allows them.
|
||||
|
||||
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.
|
||||
"""
|
||||
config = service.builtin_mcp_config(_builtin_definition(connection), writes=writes)
|
||||
if config is None:
|
||||
raise ValueError("This connector has no MCP server")
|
||||
return _classified(connection, config, _discover(user_id, connection, {"config": config}))
|
||||
|
||||
|
||||
def _rediscover(user_id: str, connection: dict, tool: dict, config: dict) -> tuple[set, set, dict]:
|
||||
"""Re-read one MCP tool's actions from ``config``'s server and store both.
|
||||
|
||||
Actions that still exist keep the permissions the user chose and the
|
||||
fixed values of parameters they still have; new ones get the defaults
|
||||
from their annotations (reads always allowed, writes need approval);
|
||||
removed ones disappear.
|
||||
|
||||
Returns:
|
||||
``(added, removed, tool)``: action names, and the serialized tool.
|
||||
"""
|
||||
discovered = _classified(connection, config, _discover(user_id, connection, {**tool, "config": config}))
|
||||
fresh = apply_default_permissions("mcp_tool", service._transform_actions(discovered))
|
||||
previous = {a.get("name"): a for a in (tool.get("actions") or []) if isinstance(a, dict)}
|
||||
merged = []
|
||||
for action in fresh:
|
||||
old = previous.get(action.get("name"))
|
||||
if old is not None:
|
||||
action = {
|
||||
**carry_pins(old, action),
|
||||
"active": old.get("active", True),
|
||||
"require_approval": bool(old.get("require_approval")),
|
||||
}
|
||||
merged.append(action)
|
||||
names = {a.get("name") for a in merged}
|
||||
fields = {"actions": merged}
|
||||
if config != (tool.get("config") or {}):
|
||||
fields["config"] = config
|
||||
with db_session() as conn:
|
||||
repo = UserToolsRepository(conn)
|
||||
repo.update(str(tool["id"]), user_id, fields)
|
||||
serialized = service.serialize_tool(repo.get_any(str(tool["id"]), user_id))
|
||||
return names - set(previous), set(previous) - names, serialized
|
||||
|
||||
|
||||
def refresh_mcp_tools(user_id: str, connection: dict) -> dict:
|
||||
"""Re-read the actions of every MCP tool on a connection.
|
||||
|
||||
A built-in connector's tool is re-read from its own endpoint: the write
|
||||
one while it is set up for writes and an admin allows them, the read-only
|
||||
one otherwise, which it then keeps.
|
||||
|
||||
Returns:
|
||||
``{"added": [...], "removed": [...], "tools": [...]}``.
|
||||
"""
|
||||
with db_readonly() as conn:
|
||||
tools = [t for t in ConnectorSessionsRepository(conn).list_tools(str(connection["id"]))
|
||||
if t.get("name") == "mcp_tool"]
|
||||
policies = service.load_policies(conn) if tools else {}
|
||||
definition = _builtin_definition(connection)
|
||||
added: set[str] = set()
|
||||
removed: set[str] = set()
|
||||
refreshed = []
|
||||
for tool in tools:
|
||||
config = dict(tool.get("config") or {})
|
||||
if definition is not None:
|
||||
config["server_url"] = service.builtin_mcp_url(
|
||||
definition, config.get("server_url"), service.writes_allowed(policies, definition.key),
|
||||
)
|
||||
tool_added, tool_removed, serialized = _rediscover(user_id, connection, tool, config)
|
||||
added |= tool_added
|
||||
removed |= tool_removed
|
||||
refreshed.append(serialized)
|
||||
return {"added": sorted(added), "removed": sorted(removed), "tools": refreshed}
|
||||
|
||||
|
||||
def set_builtin_writes(user_id: str, connection: dict, allow: bool) -> dict:
|
||||
"""Point a connection's built-in MCP tool (GitHub's) at its write or read-only endpoint.
|
||||
|
||||
The tool's actions are re-read from the new endpoint, keeping the
|
||||
choices made for actions that exist on both.
|
||||
|
||||
Args:
|
||||
user_id: The connection's owner.
|
||||
connection: The connection row.
|
||||
allow: Let agents make changes (the write endpoint) or only read.
|
||||
|
||||
Returns:
|
||||
``{"added", "removed", "tools", "writes"}``.
|
||||
|
||||
Raises:
|
||||
ValueError: The connector offers no write access.
|
||||
NoMcpTool: The connection has no such tool yet.
|
||||
service.WritesForbidden: ``allow`` while an admin forbids writes.
|
||||
service.ConnectionUnavailable: The connection needs reconnecting.
|
||||
"""
|
||||
definition = _builtin_definition(connection)
|
||||
if definition is None or not definition.mcp_write_url:
|
||||
raise ValueError("This connector has no write access to turn on")
|
||||
with db_readonly() as conn:
|
||||
tools = [t for t in ConnectorSessionsRepository(conn).list_tools(str(connection["id"]))
|
||||
if t.get("name") == "mcp_tool"]
|
||||
allowed = service.writes_allowed(service.load_policies(conn), definition.key)
|
||||
if not tools:
|
||||
raise NoMcpTool("Add this connection's tools first")
|
||||
if allow and not allowed:
|
||||
raise service.WritesForbidden(f"An admin turned off changes through {definition.name}")
|
||||
added: set[str] = set()
|
||||
removed: set[str] = set()
|
||||
refreshed = []
|
||||
for tool in tools:
|
||||
config = {**(tool.get("config") or {}), **service.builtin_mcp_config(definition, writes=allow)}
|
||||
tool_added, tool_removed, serialized = _rediscover(user_id, connection, tool, config)
|
||||
added |= tool_added
|
||||
removed |= tool_removed
|
||||
refreshed.append(serialized)
|
||||
return {"added": sorted(added), "removed": sorted(removed), "tools": refreshed, "writes": allow}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Calling a server's tools from the server side
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
SIGN_IN_EXPIRED = "The sign-in expired and could not be renewed. Reconnect to continue."
|
||||
_TRANSIENT_WORDS = ("rate limit", "ratelimit", "too many requests", "temporarily", "timed out", "timeout")
|
||||
|
||||
|
||||
class MCPToolError(Exception):
|
||||
"""An MCP tool answered with an error (a missing issue, a refused argument)."""
|
||||
|
||||
|
||||
class MCPSession:
|
||||
"""One open MCP session: list the server's tools and call them.
|
||||
|
||||
Made by :func:`run_connection_session`. Every call goes over the same
|
||||
session, so a sync of hundreds of items signs in once.
|
||||
"""
|
||||
|
||||
def __init__(self, client: Any):
|
||||
self._client = client
|
||||
self._schemas: Optional[dict[str, dict]] = None
|
||||
|
||||
async def input_schema(self, name: str) -> Optional[dict]:
|
||||
"""The input schema of tool ``name`` (``{}`` when it has none), or None when there is no such tool."""
|
||||
if self._schemas is None:
|
||||
tools = await self._client.list_tools()
|
||||
self._schemas = {tool.name: (getattr(tool, "inputSchema", None) or {}) for tool in tools}
|
||||
return self._schemas.get(name)
|
||||
|
||||
async def call(self, name: str, arguments: dict) -> Any:
|
||||
"""Call tool ``name`` and return what it answered.
|
||||
|
||||
Returns:
|
||||
Its structured content when it sends some, else its text parsed
|
||||
as JSON, else the text itself.
|
||||
|
||||
Raises:
|
||||
MCPToolError: The tool answered with an error.
|
||||
service.TransientConnectionError: The error says to try again later.
|
||||
"""
|
||||
result = await self._client.call_tool(name, arguments, raise_on_error=False)
|
||||
if getattr(result, "is_error", False):
|
||||
message = _result_text(result) or "the tool failed"
|
||||
if any(word in message.lower() for word in _TRANSIENT_WORDS):
|
||||
raise service.TransientConnectionError(f"{name}: {message}")
|
||||
raise MCPToolError(f"{name}: {message}")
|
||||
return tool_result_data(result)
|
||||
|
||||
|
||||
def _result_text(result: Any) -> str:
|
||||
blocks = getattr(result, "content", None) or []
|
||||
return "\n".join(str(block.text) for block in blocks if getattr(block, "text", None)).strip()
|
||||
|
||||
|
||||
def tool_result_data(result: Any) -> Any:
|
||||
"""What an MCP tool answered: its structured content, its text parsed as JSON, or the text."""
|
||||
structured = getattr(result, "structured_content", None)
|
||||
if isinstance(structured, dict) and structured:
|
||||
# A tool without an output schema has other values wrapped as ``{"result": ...}``.
|
||||
return structured["result"] if set(structured) == {"result"} else structured
|
||||
text = _result_text(result)
|
||||
try:
|
||||
return json.loads(text)
|
||||
except ValueError:
|
||||
return text
|
||||
|
||||
|
||||
def _client_for(connection: dict, server_url: str, timeout: float) -> Any:
|
||||
"""An MCP client for ``server_url`` that signs in with ``connection``'s stored OAuth tokens.
|
||||
|
||||
The tokens are renewed with their refresh token when they expire; when
|
||||
that fails the client raises ``MCPReauthorizationRequired`` instead of
|
||||
starting an interactive sign-in.
|
||||
"""
|
||||
from fastmcp import Client
|
||||
|
||||
from docsgpt.agents.tools.mcp_tool import MCPTool, NonInteractiveOAuth
|
||||
|
||||
tool = MCPTool(
|
||||
{"server_url": server_url, "auth_type": "oauth", "transport_type": "http", "timeout": timeout},
|
||||
connection["user_id"],
|
||||
)
|
||||
auth = NonInteractiveOAuth(
|
||||
mcp_url=tool.server_url,
|
||||
redirect_uri=tool.redirect_uri,
|
||||
user_id=connection["user_id"],
|
||||
connection_id=str(connection["id"]),
|
||||
)
|
||||
return Client(tool._create_transport(), auth=auth, timeout=timeout)
|
||||
|
||||
|
||||
def _errors(exc: BaseException):
|
||||
"""``exc`` and every exception it wraps: causes, contexts and group members."""
|
||||
seen: set[int] = set()
|
||||
pending: list[Optional[BaseException]] = [exc]
|
||||
while pending:
|
||||
current = pending.pop()
|
||||
if current is None or id(current) in seen:
|
||||
continue
|
||||
seen.add(id(current))
|
||||
yield current
|
||||
pending.extend(getattr(current, "exceptions", None) or ())
|
||||
pending.extend((current.__cause__, current.__context__))
|
||||
|
||||
|
||||
def _needs_sign_in(exc: BaseException) -> bool:
|
||||
from docsgpt.agents.tools.mcp_tool import MCPReauthorizationRequired
|
||||
|
||||
return any(
|
||||
isinstance(error, MCPReauthorizationRequired) or "oauth session expired" in str(error).lower()
|
||||
for error in _errors(exc)
|
||||
)
|
||||
|
||||
|
||||
def _is_transient(exc: BaseException) -> bool:
|
||||
import httpx
|
||||
|
||||
return any(
|
||||
isinstance(error, (ConnectionError, TimeoutError, httpx.TransportError)) for error in _errors(exc)
|
||||
)
|
||||
|
||||
|
||||
def _run_coroutine(coroutine: Awaitable[T]) -> T:
|
||||
"""Run ``coroutine`` to the end, on a thread of its own when this one already runs an event loop."""
|
||||
try:
|
||||
asyncio.get_running_loop()
|
||||
except RuntimeError:
|
||||
return asyncio.run(coroutine)
|
||||
with concurrent.futures.ThreadPoolExecutor(max_workers=1) as pool:
|
||||
return pool.submit(asyncio.run, coroutine).result()
|
||||
|
||||
|
||||
def run_connection_session(
|
||||
connection: dict,
|
||||
server_url: str,
|
||||
work: Callable[[MCPSession], Awaitable[T]],
|
||||
*,
|
||||
timeout: float = 60.0,
|
||||
) -> T:
|
||||
"""Open one MCP session signed in with ``connection`` and run ``work`` in it.
|
||||
|
||||
Args:
|
||||
connection: An MCP OAuth connection (a Linear sign-in).
|
||||
server_url: The MCP endpoint. It must be the server the connection
|
||||
signed in to, so its tokens never go anywhere else.
|
||||
work: Async function given the open :class:`MCPSession`.
|
||||
timeout: Seconds allowed for each request.
|
||||
|
||||
Returns:
|
||||
What ``work`` returned.
|
||||
|
||||
Raises:
|
||||
ValueError: ``server_url`` is not the connection's server.
|
||||
service.ConnectionUnavailable: The connection needs reconnecting,
|
||||
also when its sign-in expired and could not be renewed (it is
|
||||
then flagged, which pauses its sources).
|
||||
service.TransientConnectionError: The server could not be reached,
|
||||
or asked to slow down.
|
||||
MCPToolError: A tool answered with an error.
|
||||
"""
|
||||
connection_id = str(connection["id"])
|
||||
if catalog.base_url(connection.get("server_url")) != catalog.base_url(server_url):
|
||||
raise ValueError("A connection's sign-in only goes to its own server")
|
||||
if service.normalize_status(connection) != service.STATUS_CONNECTED:
|
||||
raise service.ConnectionUnavailable("Reconnect to continue", connection_id=connection_id)
|
||||
|
||||
async def session_work() -> T:
|
||||
async with _client_for(connection, server_url, timeout) as client:
|
||||
return await work(MCPSession(client))
|
||||
|
||||
try:
|
||||
return _run_coroutine(session_work())
|
||||
except (MCPToolError, service.TransientConnectionError, service.ConnectionUnavailable):
|
||||
raise
|
||||
except Exception as exc:
|
||||
if _needs_sign_in(exc):
|
||||
service.mark_reconnect_needed(connection_id, SIGN_IN_EXPIRED)
|
||||
raise service.ConnectionUnavailable(SIGN_IN_EXPIRED, connection_id=connection_id) from exc
|
||||
if _is_transient(exc):
|
||||
raise service.TransientConnectionError(f"The MCP server did not answer: {type(exc).__name__}") from exc
|
||||
raise
|
||||
@@ -0,0 +1,179 @@
|
||||
"""Read / write classification and permissions for connection-backed tools.
|
||||
|
||||
Every action a connection's tool offers is either a *read* (it only looks
|
||||
something up) or a *write* (it changes or sends something). Reads default to
|
||||
"Always allow"; writes default to "Needs approval". The permission is stored
|
||||
on the action itself: ``active`` off means "Off", ``require_approval`` means
|
||||
"Needs approval", neither means "Always allow".
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
from typing import Optional
|
||||
|
||||
ACCESS_READ = "read"
|
||||
ACCESS_WRITE = "write"
|
||||
|
||||
PERMISSION_ALWAYS = "always"
|
||||
PERMISSION_ASK = "ask"
|
||||
PERMISSION_OFF = "off"
|
||||
PERMISSIONS = (PERMISSION_ALWAYS, PERMISSION_ASK, PERMISSION_OFF)
|
||||
|
||||
_READ_WORDS = frozenset(
|
||||
("search", "query", "find", "list", "get", "read", "fetch", "lookup", "describe", "retrieve")
|
||||
)
|
||||
# A name holding any of these is a write even when it also holds a read word
|
||||
# (``get_or_create_page``, ``search_and_replace``).
|
||||
_WRITE_WORDS = frozenset((
|
||||
"create", "update", "delete", "remove", "set", "send", "post", "put", "patch", "write", "add", "insert",
|
||||
"upsert", "append", "replace", "edit", "modify", "rename", "move", "copy", "upload", "save", "submit",
|
||||
"publish", "share", "invite", "assign", "archive", "restore", "reset", "clear", "purge", "drop", "execute",
|
||||
"run", "invoke", "trigger", "start", "stop", "cancel", "close", "merge", "approve", "reject", "reply",
|
||||
"comment", "mark", "enable", "disable", "grant", "revoke", "transfer", "import", "sync", "lock", "unlock",
|
||||
))
|
||||
_CAMEL_BOUNDARY = re.compile(r"(?<=[a-z0-9])(?=[A-Z])|(?<=[A-Z])(?=[A-Z][a-z])")
|
||||
_NON_WORD = re.compile(r"[^A-Za-z0-9]+")
|
||||
|
||||
|
||||
def _name_words(name: str) -> set[str]:
|
||||
"""The lower-cased words of an action name (``snake``, ``kebab``, ``camelCase``)."""
|
||||
return {word.lower() for word in _NON_WORD.split(_CAMEL_BOUNDARY.sub(" ", name)) if word}
|
||||
|
||||
|
||||
def action_access(tool_name: Optional[str], action: dict) -> str:
|
||||
"""Whether ``action`` reads or writes.
|
||||
|
||||
Order of evidence: an explicit ``access`` on the action metadata, MCP tool
|
||||
annotations (``readOnlyHint`` / ``destructiveHint``), the HTTP method of an
|
||||
API tool action, then the action's name: a read only when one of its
|
||||
words is a read verb and none is a write verb.
|
||||
|
||||
Args:
|
||||
tool_name: The ``user_tools`` name the action belongs to.
|
||||
action: One entry of the tool's ``actions``.
|
||||
|
||||
Returns:
|
||||
``read`` or ``write``.
|
||||
"""
|
||||
access = action.get("access")
|
||||
if access in (ACCESS_READ, ACCESS_WRITE):
|
||||
return access
|
||||
annotations = action.get("annotations") or {}
|
||||
if isinstance(annotations, dict):
|
||||
if annotations.get("readOnlyHint") is True:
|
||||
return ACCESS_READ
|
||||
if annotations.get("destructiveHint") is True or annotations.get("readOnlyHint") is False:
|
||||
return ACCESS_WRITE
|
||||
method = (action.get("method") or "").upper()
|
||||
if tool_name == "api_tool" and method:
|
||||
return ACCESS_READ if method in ("GET", "HEAD", "OPTIONS") else ACCESS_WRITE
|
||||
words = _name_words(action.get("name") or "")
|
||||
return ACCESS_READ if words & _READ_WORDS and not words & _WRITE_WORDS else ACCESS_WRITE
|
||||
|
||||
|
||||
def action_permission(action: dict) -> str:
|
||||
"""The action's permission: ``always``, ``ask`` or ``off``."""
|
||||
if action.get("active") is False:
|
||||
return PERMISSION_OFF
|
||||
if action.get("require_approval"):
|
||||
return PERMISSION_ASK
|
||||
return PERMISSION_ALWAYS
|
||||
|
||||
|
||||
def apply_permission(action: dict, permission: str) -> dict:
|
||||
"""Return ``action`` with ``permission`` written onto its flags."""
|
||||
if permission not in PERMISSIONS:
|
||||
raise ValueError(f"Unknown permission: {permission}")
|
||||
updated = dict(action)
|
||||
updated["active"] = permission != PERMISSION_OFF
|
||||
updated["require_approval"] = permission == PERMISSION_ASK
|
||||
return updated
|
||||
|
||||
|
||||
def apply_default_permissions(tool_name: Optional[str], actions: list[dict]) -> list[dict]:
|
||||
"""Stamp ``access`` on each action and default writes to needing approval."""
|
||||
stamped = []
|
||||
for action in actions:
|
||||
access = action_access(tool_name, action)
|
||||
updated = {**action, "access": access}
|
||||
if access == ACCESS_WRITE:
|
||||
updated["require_approval"] = True
|
||||
stamped.append(updated)
|
||||
return stamped
|
||||
|
||||
|
||||
def tool_actions(tool: dict) -> list[dict]:
|
||||
"""A tool row's actions, each carrying its ``name``.
|
||||
|
||||
An API tool keeps its actions under ``config["actions"]`` keyed by name;
|
||||
every other tool lists them in ``actions``.
|
||||
"""
|
||||
if tool.get("name") == "api_tool":
|
||||
stored = (tool.get("config") or {}).get("actions") or {}
|
||||
return [{**(action or {}), "name": name} for name, action in stored.items()]
|
||||
return [action for action in tool.get("actions") or [] if isinstance(action, dict)]
|
||||
|
||||
|
||||
# Where an API tool keeps the values it sends: header and query values are
|
||||
# the owner's (sealed per action, flagged ``has_value``; legacy rows hold them
|
||||
# in plain ``value``). Mirrors ``tool_executor.API_TOOL_SECRET_SECTIONS``.
|
||||
_API_TOOL_SECRET_SECTIONS = ("headers", "query_params")
|
||||
|
||||
|
||||
def _api_action_sends_credentials(action: dict) -> bool:
|
||||
"""Whether an API tool action sends a header or query value the owner stored."""
|
||||
for section in _API_TOOL_SECRET_SECTIONS:
|
||||
block = action.get(section)
|
||||
props = block.get("properties") if isinstance(block, dict) else None
|
||||
for spec in (props or {}).values():
|
||||
if isinstance(spec, dict) and (spec.get("has_value") or spec.get("value") not in (None, "")):
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
def holds_owner_credentials(tool: dict, action_name: Optional[str] = None) -> bool:
|
||||
"""Whether ``tool`` (or its ``action_name``) acts with credentials its owner stored.
|
||||
|
||||
A connection's account, a stored secret, an MCP server the owner signed
|
||||
in to, or an API tool action that sends a header or query value the
|
||||
owner saved (a key, a token). Anyone running it acts as the owner there,
|
||||
whoever they are. An API tool action that sends nothing stored does not.
|
||||
|
||||
Args:
|
||||
tool: A ``user_tools`` row.
|
||||
action_name: One action to judge; None judges the tool as a whole.
|
||||
|
||||
Returns:
|
||||
True when the tool (or that action) runs on the owner's credentials.
|
||||
"""
|
||||
config = tool.get("config") or {}
|
||||
if tool.get("connection_id") or config.get("encrypted_credentials"):
|
||||
return True
|
||||
if tool.get("name") == "api_tool":
|
||||
actions = config.get("actions") or {}
|
||||
if action_name is not None:
|
||||
return _api_action_sends_credentials(actions.get(action_name) or {})
|
||||
return any(_api_action_sends_credentials(action or {}) for action in actions.values())
|
||||
return tool.get("name") == "mcp_tool" and (config.get("auth_type") or "none") != "none"
|
||||
|
||||
|
||||
def owner_credential_writes(tool: dict) -> list[str]:
|
||||
"""Names of the tool's write actions that run on its owner's credentials.
|
||||
|
||||
These are what someone who can't approve for the owner (an API-key or
|
||||
widget caller, a public-link user) may run only when the owner allows
|
||||
them in the agent's API write allowlist.
|
||||
|
||||
Args:
|
||||
tool: A ``user_tools`` row.
|
||||
|
||||
Returns:
|
||||
Action names, empty when the tool holds no owner credentials.
|
||||
"""
|
||||
return [
|
||||
action["name"] for action in tool_actions(tool)
|
||||
if action.get("name") and action.get("active") is not False
|
||||
and action_access(tool.get("name"), action) == ACCESS_WRITE
|
||||
and holds_owner_credentials(tool, action["name"])
|
||||
]
|
||||
@@ -0,0 +1,77 @@
|
||||
# Curated remote MCP servers shown as Connectors catalog cards.
|
||||
#
|
||||
# Inclusion rules: the vendor operates the server, it has a stable public URL,
|
||||
# and it signs in with OAuth plus dynamic client registration (or needs no
|
||||
# auth), so connecting takes no admin setup. Each entry was checked for OAuth
|
||||
# authorization-server metadata with a registration endpoint and for a 401
|
||||
# challenge on the MCP endpoint. Keep the list short; anything else is a
|
||||
# custom connector.
|
||||
#
|
||||
# Keys look like ``mcp:<name>``. The frontend reads ``settings.connectors
|
||||
# .descriptions.mcp_<name>`` for a translated description.
|
||||
#
|
||||
# A preset with ``sync_ingestor`` also syncs into Knowledge through its own
|
||||
# MCP server, signed in with the same connection (Linear's issues).
|
||||
|
||||
- key: mcp:notion
|
||||
name: Notion
|
||||
description: Search, read and update Notion pages and databases.
|
||||
icon: notion
|
||||
category: knowledge
|
||||
mcp_url: https://mcp.notion.com/mcp
|
||||
auth_kind: mcp_oauth
|
||||
capabilities: [read, write]
|
||||
docs_url: https://developers.notion.com/docs/mcp
|
||||
|
||||
- key: mcp:linear
|
||||
name: Linear
|
||||
description: Sync issues and documents into Knowledge, and let agents find, create and update issues.
|
||||
icon: linear
|
||||
category: projects
|
||||
mcp_url: https://mcp.linear.app/mcp
|
||||
auth_kind: mcp_oauth
|
||||
capabilities: [sync, read, write]
|
||||
sync_ingestor: linear
|
||||
docs_url: https://linear.app/docs/mcp
|
||||
|
||||
- key: mcp:atlassian
|
||||
name: Jira & Confluence
|
||||
description: Search and update Jira issues and Confluence pages.
|
||||
icon: atlassian
|
||||
category: projects
|
||||
# Shown on the Confluence card, next to syncing pages into Knowledge.
|
||||
part_of: confluence
|
||||
mcp_url: https://mcp.atlassian.com/v1/mcp
|
||||
auth_kind: mcp_oauth
|
||||
capabilities: [read, write]
|
||||
docs_url: https://support.atlassian.com/rovo/docs/getting-started-with-the-atlassian-remote-mcp-server/
|
||||
|
||||
- key: mcp:sentry
|
||||
name: Sentry
|
||||
description: Look up Sentry issues, events and releases.
|
||||
icon: sentry
|
||||
category: dev
|
||||
mcp_url: https://mcp.sentry.dev/mcp
|
||||
auth_kind: mcp_oauth
|
||||
capabilities: [read, write]
|
||||
docs_url: https://docs.sentry.io/product/sentry-mcp/
|
||||
|
||||
- key: mcp:asana
|
||||
name: Asana
|
||||
description: Find and update Asana tasks and projects.
|
||||
icon: asana
|
||||
category: projects
|
||||
mcp_url: https://mcp.asana.com/v2/mcp
|
||||
auth_kind: mcp_oauth
|
||||
capabilities: [read, write]
|
||||
docs_url: https://developers.asana.com/docs/using-asanas-mcp-server
|
||||
|
||||
- key: mcp:stripe
|
||||
name: Stripe
|
||||
description: Look up customers, payments and subscriptions in Stripe.
|
||||
icon: stripe
|
||||
category: business
|
||||
mcp_url: https://mcp.stripe.com
|
||||
auth_kind: mcp_oauth
|
||||
capabilities: [read, write]
|
||||
docs_url: https://docs.stripe.com/mcp
|
||||
@@ -0,0 +1,279 @@
|
||||
"""Which connection a shared tool or source uses at runtime.
|
||||
|
||||
A resource that points at a connection runs either with its owner's account
|
||||
(``owner`` mode, the default for new shares) or with the invoking member's
|
||||
own account for the same service (``member`` mode). Resolution never returns
|
||||
credentials; callers read them from the resolved row through
|
||||
``docsgpt.connectors.service``.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime, timezone
|
||||
import logging
|
||||
from dataclasses import dataclass
|
||||
from typing import Optional
|
||||
|
||||
from docsgpt.connectors import catalog, service
|
||||
from docsgpt.storage.db.repositories.connector_sessions import ConnectorSessionsRepository
|
||||
from docsgpt.storage.db.session import db_readonly
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
MODE_OWNER = "owner"
|
||||
MODE_MEMBER = "member"
|
||||
|
||||
# Why a connection-backed tool can't run (``connection_stop_reason``).
|
||||
CONNECTION_NEEDS_RECONNECT = "connection_needs_reconnect"
|
||||
CONNECTION_REMOVED = "connection_removed"
|
||||
CONNECTOR_DISABLED = "connector_disabled"
|
||||
|
||||
# Tool config key ``remove_connection`` sets when it keeps a connection's
|
||||
# tools: the connector key, or True when the connection had none, so they can
|
||||
# say what they lost. Only the server writes it (see
|
||||
# :func:`carry_removed_connection`).
|
||||
REMOVED_CONNECTION_KEY = "removed_connection"
|
||||
|
||||
|
||||
def carry_removed_connection(new_config: dict, stored_config: Optional[dict]) -> dict:
|
||||
"""``new_config`` with the stored removed-connection note, never a client's.
|
||||
|
||||
A config save must neither fake the note nor clear it; every path that
|
||||
writes a tool config from a request passes it through here.
|
||||
|
||||
Args:
|
||||
new_config: The config about to be stored.
|
||||
stored_config: The tool's stored config (None or ``{}`` when creating it).
|
||||
|
||||
Returns:
|
||||
dict: A copy of ``new_config`` carrying the stored note, if any.
|
||||
"""
|
||||
out = {k: v for k, v in (new_config or {}).items() if k != REMOVED_CONNECTION_KEY}
|
||||
marker = (stored_config or {}).get(REMOVED_CONNECTION_KEY)
|
||||
if marker:
|
||||
out[REMOVED_CONNECTION_KEY] = marker
|
||||
return out
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ResolvedConnection:
|
||||
"""The connection a call runs with, or why there is none.
|
||||
|
||||
Attributes:
|
||||
row: The ``connector_sessions`` row, None when missing.
|
||||
available: Whether it can be used right now.
|
||||
connector_key: Catalog key of the service.
|
||||
connector_name: Name shown to the user ("Notion", or a custom label).
|
||||
delegated: The row belongs to someone other than the invoker.
|
||||
writes_allowed: Whether an admin lets agents make changes through a
|
||||
connector that offers them as an opt-in (GitHub); True elsewhere.
|
||||
enabled: Whether the connector is switched on (an admin can turn it off).
|
||||
mode: Whose account the resource runs with, after any mode an admin
|
||||
forces: :data:`MODE_OWNER` or :data:`MODE_MEMBER`.
|
||||
"""
|
||||
|
||||
row: Optional[dict]
|
||||
available: bool
|
||||
connector_key: Optional[str]
|
||||
connector_name: Optional[str]
|
||||
delegated: bool = False
|
||||
writes_allowed: bool = True
|
||||
enabled: bool = True
|
||||
mode: str = MODE_OWNER
|
||||
|
||||
@property
|
||||
def connection_id(self) -> Optional[str]:
|
||||
return str(self.row["id"]) if self.row else None
|
||||
|
||||
|
||||
def _name_for(row: Optional[dict], fallback_key: Optional[str]) -> Optional[str]:
|
||||
if row is not None:
|
||||
return service.serialize_connection(row)["name"]
|
||||
definition = catalog.get_definition(fallback_key)
|
||||
return definition.name if definition else None
|
||||
|
||||
|
||||
def resolve_connection(
|
||||
resource: dict,
|
||||
invoker_user_id: Optional[str],
|
||||
*,
|
||||
conn=None,
|
||||
policies: Optional[dict] = None,
|
||||
) -> Optional[ResolvedConnection]:
|
||||
"""Pick the connection a tool or source uses for ``invoker_user_id``.
|
||||
|
||||
``owner`` mode uses ``resource.connection_id``. ``member`` mode uses the
|
||||
invoker's own connection to the same service (and, for MCP, the same
|
||||
server), falling back to the owner's when the invoker is the owner.
|
||||
|
||||
Args:
|
||||
resource: A ``user_tools`` or ``sources`` row.
|
||||
invoker_user_id: Who is running it.
|
||||
conn: An open connection to reuse; a read-only one is opened when None.
|
||||
policies: Connector policies already loaded with ``service.load_policies``,
|
||||
so a caller resolving many resources loads them once.
|
||||
|
||||
Returns:
|
||||
None when the resource has no connection at all; otherwise the
|
||||
resolution, possibly with ``available=False``.
|
||||
"""
|
||||
if not resource.get("connection_id"):
|
||||
return None
|
||||
if conn is None:
|
||||
with db_readonly() as own_conn:
|
||||
return _resolve(own_conn, resource, invoker_user_id, policies)
|
||||
return _resolve(conn, resource, invoker_user_id, policies)
|
||||
|
||||
|
||||
def effective_credential_mode(resource: dict, policies: dict, connector_key: Optional[str]) -> str:
|
||||
"""Whose account a connection-backed resource runs with: ``owner`` or ``member``.
|
||||
|
||||
The resource's own ``credential_mode``, unless an admin forces one mode
|
||||
for every share of its connector.
|
||||
|
||||
Args:
|
||||
resource: A ``user_tools`` or ``sources`` row.
|
||||
policies: Connector policies from ``service.load_policies``.
|
||||
connector_key: Catalog key of the resource's connection.
|
||||
|
||||
Returns:
|
||||
:data:`MODE_OWNER` or :data:`MODE_MEMBER`.
|
||||
"""
|
||||
policy = (policies.get(connector_key) or {}) if connector_key else {}
|
||||
if policy.get("credential_mode") in (MODE_OWNER, MODE_MEMBER):
|
||||
return policy["credential_mode"]
|
||||
return MODE_MEMBER if resource.get("credential_mode") == MODE_MEMBER else MODE_OWNER
|
||||
|
||||
|
||||
def _resolve(conn, resource: dict, invoker_user_id: Optional[str], policies: Optional[dict]) -> ResolvedConnection:
|
||||
connection_id = resource.get("connection_id")
|
||||
owner = resource.get("user_id")
|
||||
repo = ConnectorSessionsRepository(conn)
|
||||
owned = repo.get(str(connection_id))
|
||||
owned_key = catalog.connector_key_for_row(owned) if owned else None
|
||||
if policies is None:
|
||||
policies = service.load_policies(conn)
|
||||
mode = effective_credential_mode(resource, policies, owned_key)
|
||||
if owned is not None and owner and owned.get("user_id") != owner:
|
||||
# A resource may only point at its own owner's connection.
|
||||
logger.warning(
|
||||
"resource %s points at a connection it does not own", resource.get("id"),
|
||||
)
|
||||
owned = None
|
||||
row = owned
|
||||
if mode == MODE_MEMBER and invoker_user_id and invoker_user_id != owner:
|
||||
row = _member_connection(repo, owned, invoker_user_id)
|
||||
key = catalog.connector_key_for_row(row or owned or {})
|
||||
enabled = service.connector_is_enabled(policies, key)
|
||||
available = (
|
||||
row is not None
|
||||
and service.normalize_status(row) == service.STATUS_CONNECTED
|
||||
and enabled
|
||||
)
|
||||
return ResolvedConnection(
|
||||
row=row,
|
||||
available=available,
|
||||
connector_key=key,
|
||||
connector_name=_name_for(row or owned, key),
|
||||
delegated=bool(row and invoker_user_id and row.get("user_id") != invoker_user_id),
|
||||
writes_allowed=_writes_allowed(policies, key),
|
||||
enabled=enabled,
|
||||
mode=mode,
|
||||
)
|
||||
|
||||
|
||||
def connection_stop_reason(tool: dict, resolved: Optional[ResolvedConnection]) -> Optional[str]:
|
||||
"""Why an owner-mode tool's connection keeps it from running, or None.
|
||||
|
||||
Only the connection the owner's account runs on is judged. A member-mode
|
||||
tool runs on each caller's own account, so the owner's account says
|
||||
nothing about whether it runs; only an admin turning the service off
|
||||
stops it for everyone.
|
||||
|
||||
Args:
|
||||
tool: The ``user_tools`` row.
|
||||
resolved: What :func:`resolve_connection` returned for it, resolved
|
||||
for the tool's holder's owner.
|
||||
|
||||
Returns:
|
||||
:data:`CONNECTION_REMOVED` when its connection was removed and the
|
||||
tool kept (``remove_connection`` notes it on the tool) with no
|
||||
credentials of its own, or points at
|
||||
one it may not use; :data:`CONNECTOR_DISABLED` when an admin turned
|
||||
the service off; :data:`CONNECTION_NEEDS_RECONNECT` when the account
|
||||
must sign in again; else None.
|
||||
"""
|
||||
if resolved is None:
|
||||
# Only a tool that had a connection lost it: a tool that never had
|
||||
# one (a tokenless ntfy, a legacy tool) runs on its own config, and
|
||||
# so does a kept one its owner gave credentials of its own since.
|
||||
config = tool.get("config") or {}
|
||||
removed = config.get(REMOVED_CONNECTION_KEY) and not config.get("encrypted_credentials")
|
||||
return CONNECTION_REMOVED if removed and not tool.get("connection_id") else None
|
||||
if not resolved.enabled:
|
||||
return CONNECTOR_DISABLED
|
||||
if resolved.mode == MODE_MEMBER:
|
||||
return None
|
||||
if resolved.row is None:
|
||||
return CONNECTION_REMOVED
|
||||
if not resolved.available:
|
||||
return CONNECTION_NEEDS_RECONNECT
|
||||
return None
|
||||
|
||||
|
||||
def _writes_allowed(policies: dict, key: Optional[str]) -> bool:
|
||||
definition = catalog.get_definition(key) if key else None
|
||||
if definition is None or not definition.mcp_write_url:
|
||||
return True
|
||||
return service.writes_allowed(policies, key)
|
||||
|
||||
|
||||
def _member_connection(repo: ConnectorSessionsRepository, owned: Optional[dict], invoker: str) -> Optional[dict]:
|
||||
"""The invoker's own connection to the service the owner's connection is for.
|
||||
|
||||
A member with several connected accounts of that service gets the one
|
||||
most recently used or connected, whichever is later: the account they
|
||||
are working in, or the one a "Connect to continue" just added (never
|
||||
used yet). Ties go to the last used, then the last updated.
|
||||
"""
|
||||
if owned is None:
|
||||
return None
|
||||
candidates = [
|
||||
row for row in repo.list_for_user(invoker)
|
||||
if row.get("provider") == owned.get("provider")
|
||||
and (row.get("server_url") or "") == (owned.get("server_url") or "")
|
||||
and service.normalize_status(row) == service.STATUS_CONNECTED
|
||||
]
|
||||
if not candidates:
|
||||
return None
|
||||
|
||||
def when(value) -> datetime:
|
||||
if isinstance(value, str):
|
||||
value = datetime.fromisoformat(value)
|
||||
if not isinstance(value, datetime):
|
||||
return datetime.min.replace(tzinfo=timezone.utc)
|
||||
return value if value.tzinfo else value.replace(tzinfo=timezone.utc)
|
||||
|
||||
def recency(row: dict) -> tuple:
|
||||
used, updated, created = (when(row.get(field)) for field in ("last_used_at", "updated_at", "created_at"))
|
||||
return max(used, created), used, updated, created
|
||||
|
||||
return max(candidates, key=recency)
|
||||
|
||||
|
||||
def audit_delegation(resolved: ResolvedConnection, *, invoker: Optional[str], resource_type: str,
|
||||
resource_id: Optional[str], agent_id: Optional[str] = None) -> None:
|
||||
"""Log a call that runs with someone else's account (``owner`` mode)."""
|
||||
if not resolved.delegated or resolved.row is None:
|
||||
return
|
||||
logger.info(
|
||||
"tool_credential_delegation",
|
||||
extra={
|
||||
"invoker": invoker,
|
||||
"tool_owner": resolved.row.get("user_id"),
|
||||
"connection_id": resolved.connection_id,
|
||||
"resource_type": resource_type,
|
||||
"resource_id": resource_id,
|
||||
"agent_id": agent_id,
|
||||
},
|
||||
)
|
||||
File diff suppressed because it is too large.
Load diff
@@ -29,7 +29,17 @@ class AuthSettings(SettingsGroup):
|
||||
)
|
||||
ENCRYPTION_SECRET_KEY: str = Field(
|
||||
default="default-docsgpt-encryption-key",
|
||||
description="Key used to encrypt stored credentials such as tool and connector secrets.",
|
||||
description=(
|
||||
"Key used to encrypt stored credentials such as tool and connector secrets. Set your own value before "
|
||||
"connecting services on a multi-user install; the default is public."
|
||||
),
|
||||
)
|
||||
ENCRYPTION_SECRET_KEY_PREVIOUS: Optional[str] = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"Previous ENCRYPTION_SECRET_KEY, tried when a stored credential was encrypted with it. Set it while "
|
||||
"rotating the key, run `docsgpt connectors reencrypt`, then remove it."
|
||||
),
|
||||
)
|
||||
INTERNAL_KEY: Optional[str] = Field(
|
||||
default=None, description="Internal API key for worker-to-backend authentication."
|
||||
|
||||
@@ -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."
|
||||
|
||||
@@ -212,6 +212,25 @@ class AgentConfig(BaseModel):
|
||||
model_config = ConfigDict(extra="forbid")
|
||||
|
||||
guardrails: GuardrailsConfig = GuardrailsConfig()
|
||||
# Write actions on credentials the owner holds (a connected account, a
|
||||
# saved API tool key, a stored secret, an MCP sign-in) that someone who
|
||||
# can't approve for the owner may run: an API-key or widget caller, a
|
||||
# public-link user, and schedules either of them set. Any other such
|
||||
# write is refused for them. ``tool_id:action``.
|
||||
api_write_allowlist: List[str] = []
|
||||
|
||||
@field_validator("api_write_allowlist")
|
||||
@classmethod
|
||||
def _check_allowlist(cls, value: List[str]) -> List[str]:
|
||||
if len(value) > 200:
|
||||
raise ValueError("api_write_allowlist accepts at most 200 actions")
|
||||
cleaned = []
|
||||
for entry in value:
|
||||
tool_id, _, action = str(entry).partition(":")
|
||||
if not tool_id.strip() or not action.strip():
|
||||
raise ValueError("api_write_allowlist entries must be 'tool_id:action'")
|
||||
cleaned.append(f"{tool_id.strip()}:{action.strip()}")
|
||||
return sorted(set(cleaned))
|
||||
|
||||
@classmethod
|
||||
def parse(cls, raw: Optional[dict]) -> "AgentConfig":
|
||||
@@ -221,4 +240,5 @@ class AgentConfig(BaseModel):
|
||||
try:
|
||||
return cls.model_validate(raw)
|
||||
except Exception:
|
||||
# A bad allowlist falls back to none: the safe side.
|
||||
return cls(guardrails=GuardrailsConfig.parse(raw.get("guardrails")))
|
||||
@@ -1191,27 +1191,11 @@ class LLMHandler(ABC):
|
||||
)
|
||||
if hasattr(agent.tool_executor, "headless_denials"):
|
||||
agent.tool_executor.headless_denials.append(pause_info)
|
||||
from docsgpt.agents.tool_executor import (
|
||||
_mark_failed,
|
||||
_record_proposed,
|
||||
)
|
||||
from docsgpt.agents.tool_executor import journal_refused_call
|
||||
|
||||
if _record_proposed(
|
||||
pause_info["call_id"],
|
||||
pause_info["tool_name"],
|
||||
pause_info["action_name"],
|
||||
pause_info.get("arguments") or {},
|
||||
tool_id=pause_info.get("tool_id"),
|
||||
message_id=agent.tool_executor.message_id,
|
||||
user_id=agent.tool_executor.user,
|
||||
agent_id=agent.tool_executor.agent_id,
|
||||
):
|
||||
_mark_failed(
|
||||
pause_info["call_id"],
|
||||
f"headless: {deny_reason}",
|
||||
message_id=agent.tool_executor.message_id,
|
||||
user_id=agent.tool_executor.user,
|
||||
)
|
||||
journal_refused_call(
|
||||
agent.tool_executor, pause_info, f"headless: {deny_reason}"
|
||||
)
|
||||
denied_data = {
|
||||
"tool_name": pause_info["tool_name"],
|
||||
"call_id": pause_info["call_id"],
|
||||
@@ -1240,6 +1224,13 @@ class LLMHandler(ABC):
|
||||
# can wire the sticky "don't ask again" button.
|
||||
if pause_info.get("device_id"):
|
||||
pause_data["device_id"] = pause_info["device_id"]
|
||||
# What will be sent once fixed values replace the model's.
|
||||
if pause_info.get("sent_arguments") is not None:
|
||||
pause_data["sent_arguments"] = pause_info["sent_arguments"]
|
||||
# A connection-backed tool whose account needs signing in: the
|
||||
# approval card becomes a Connect card.
|
||||
if pause_info.get("connection_required"):
|
||||
pause_data["connection_required"] = pause_info["connection_required"]
|
||||
trace_unexecuted_tool_call(call, pause_data)
|
||||
yield {"type": "tool_call", "data": pause_data}
|
||||
pending_actions.append(pause_info)
|
||||
|
||||
@@ -6,7 +6,7 @@ interface for external knowledge base connectors.
|
||||
"""
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Any, Dict, List, Optional
|
||||
from typing import Any, Dict, List, Optional, Tuple
|
||||
|
||||
from docsgpt.parser.schema.base import Document
|
||||
|
||||
@@ -88,18 +88,55 @@ class BaseConnectorLoader(ABC):
|
||||
Abstract base class for connector loaders.
|
||||
|
||||
Defines the minimal interface that all connector loader
|
||||
implementations must follow.
|
||||
implementations must follow. A loader reads its OAuth tokens through
|
||||
``docsgpt.connectors.service`` from the connection it was built for,
|
||||
either directly (``connection_id``, what background sync uses) or through
|
||||
a legacy browser ``session_token`` that names the connection.
|
||||
"""
|
||||
|
||||
|
||||
connection_id: Optional[str] = None
|
||||
|
||||
@abstractmethod
|
||||
def __init__(self, session_token: str):
|
||||
def __init__(self, session_token: Optional[str] = None, *, connection_id: Optional[str] = None):
|
||||
"""
|
||||
Initialize the connector loader.
|
||||
|
||||
Args:
|
||||
session_token: Authentication session token
|
||||
session_token: Legacy browser session token naming the connection.
|
||||
connection_id: The connection to read tokens from.
|
||||
"""
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def from_connection(cls, connection_id: str) -> "BaseConnectorLoader":
|
||||
"""Build a loader that reads its tokens from ``connection_id``."""
|
||||
return cls(connection_id=connection_id)
|
||||
|
||||
def _load_token_info(
|
||||
self, session_token: Optional[str], connection_id: Optional[str],
|
||||
) -> Tuple[str, Dict[str, Any]]:
|
||||
"""Resolve the connection and return ``(connection_id, token_info)``.
|
||||
|
||||
Raises:
|
||||
ValueError: The connection is missing or needs reconnecting.
|
||||
"""
|
||||
from docsgpt.connectors import service
|
||||
|
||||
resolved = connection_id or service.connection_id_for_session_token(session_token)
|
||||
self.connection_id = resolved
|
||||
return resolved, service.get_valid_token_info(resolved)
|
||||
|
||||
def _refresh_rejected_token(self, access_token: Optional[str]) -> Dict[str, Any]:
|
||||
"""Token info after the provider answered 401 to ``access_token``.
|
||||
|
||||
Refreshes under the connection's row lock and persists the rotated
|
||||
refresh token, or returns the token another worker already renewed.
|
||||
"""
|
||||
from docsgpt.connectors import service
|
||||
|
||||
if not self.connection_id:
|
||||
raise ValueError("Loader has no connection to refresh")
|
||||
return service.get_valid_token_info(self.connection_id, rejected_access_token=access_token)
|
||||
|
||||
@abstractmethod
|
||||
def load_data(self, inputs: Dict[str, Any]) -> List[Document]:
|
||||
|
||||
@@ -6,7 +6,6 @@ from urllib.parse import urlencode
|
||||
import requests
|
||||
|
||||
from docsgpt.core.settings import settings
|
||||
from docsgpt.parser.connectors._auth_utils import session_token_fingerprint
|
||||
from docsgpt.parser.connectors.base import BaseConnectorAuth
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -152,32 +151,6 @@ class ConfluenceAuth(BaseConnectorAuth):
|
||||
except Exception:
|
||||
return True
|
||||
|
||||
def get_token_info_from_session(self, session_token: str) -> Dict[str, Any]:
|
||||
from docsgpt.storage.db.repositories.connector_sessions import (
|
||||
ConnectorSessionsRepository,
|
||||
)
|
||||
from docsgpt.storage.db.session import db_readonly
|
||||
|
||||
with db_readonly() as conn:
|
||||
session = ConnectorSessionsRepository(conn).get_by_session_token(
|
||||
session_token
|
||||
)
|
||||
if not session:
|
||||
raise ValueError(
|
||||
f"Invalid session token ({session_token_fingerprint(session_token)})"
|
||||
)
|
||||
|
||||
token_info = session.get("token_info")
|
||||
if not token_info:
|
||||
raise ValueError("Session missing token information")
|
||||
|
||||
required = ["access_token", "refresh_token", "cloud_id"]
|
||||
missing = [f for f in required if not token_info.get(f)]
|
||||
if missing:
|
||||
raise ValueError(f"Missing required token fields: {missing}")
|
||||
|
||||
return token_info
|
||||
|
||||
def sanitize_token_info(
|
||||
self, token_info: Dict[str, Any], **extra_fields
|
||||
) -> Dict[str, Any]:
|
||||
|
||||
@@ -44,12 +44,8 @@ def _retry_on_auth_failure(func):
|
||||
"Auth failure in %s, refreshing token and retrying", func.__name__
|
||||
)
|
||||
try:
|
||||
new_token_info = self.auth.refresh_access_token(self.refresh_token)
|
||||
new_token_info = self._refresh_rejected_token(self.access_token)
|
||||
self.access_token = new_token_info["access_token"]
|
||||
self.refresh_token = new_token_info.get(
|
||||
"refresh_token", self.refresh_token
|
||||
)
|
||||
self._persist_refreshed_tokens(new_token_info)
|
||||
except Exception as refresh_err:
|
||||
raise ValueError(
|
||||
f"Authentication failed and could not be refreshed: {refresh_err}"
|
||||
@@ -62,13 +58,15 @@ def _retry_on_auth_failure(func):
|
||||
|
||||
class ConfluenceLoader(BaseConnectorLoader):
|
||||
|
||||
def __init__(self, session_token: str):
|
||||
def __init__(self, session_token: Optional[str] = None, *, connection_id: Optional[str] = None):
|
||||
self.auth = ConfluenceAuth()
|
||||
self.session_token = session_token
|
||||
|
||||
token_info = self.auth.get_token_info_from_session(session_token)
|
||||
_, token_info = self._load_token_info(session_token, connection_id)
|
||||
missing = [f for f in ("access_token", "cloud_id") if not token_info.get(f)]
|
||||
if missing:
|
||||
raise ValueError(f"Missing required token fields: {missing}")
|
||||
self.access_token = token_info["access_token"]
|
||||
self.refresh_token = token_info["refresh_token"]
|
||||
self.cloud_id = token_info["cloud_id"]
|
||||
|
||||
self.base_url = API_V2.format(cloud_id=self.cloud_id)
|
||||
@@ -81,22 +79,6 @@ class ConfluenceLoader(BaseConnectorLoader):
|
||||
"Accept": "application/json",
|
||||
}
|
||||
|
||||
def _persist_refreshed_tokens(self, token_info: Dict[str, Any]) -> None:
|
||||
try:
|
||||
from docsgpt.storage.db.repositories.connector_sessions import (
|
||||
ConnectorSessionsRepository,
|
||||
)
|
||||
from docsgpt.storage.db.session import db_session
|
||||
|
||||
sanitized = self.auth.sanitize_token_info(token_info)
|
||||
with db_session() as conn:
|
||||
repo = ConnectorSessionsRepository(conn)
|
||||
session = repo.get_by_session_token(self.session_token)
|
||||
if session:
|
||||
repo.update(str(session["id"]), {"token_info": sanitized})
|
||||
except Exception as e:
|
||||
logger.warning("Failed to persist refreshed tokens: %s", e)
|
||||
|
||||
@_retry_on_auth_failure
|
||||
def load_data(self, inputs: Dict[str, Any]) -> List[Document]:
|
||||
folder_id = inputs.get("folder_id")
|
||||
|
||||
@@ -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 ""}
|
||||
@@ -8,7 +8,6 @@ from googleapiclient.discovery import build
|
||||
from googleapiclient.errors import HttpError
|
||||
|
||||
from docsgpt.core.settings import settings
|
||||
from docsgpt.parser.connectors._auth_utils import session_token_fingerprint
|
||||
from docsgpt.parser.connectors.base import BaseConnectorAuth
|
||||
|
||||
|
||||
@@ -213,39 +212,6 @@ class GoogleDriveAuth(BaseConnectorAuth):
|
||||
|
||||
return True
|
||||
|
||||
def get_token_info_from_session(self, session_token: str) -> Dict[str, Any]:
|
||||
try:
|
||||
from docsgpt.storage.db.repositories.connector_sessions import (
|
||||
ConnectorSessionsRepository,
|
||||
)
|
||||
from docsgpt.storage.db.session import db_readonly
|
||||
|
||||
with db_readonly() as conn:
|
||||
session = ConnectorSessionsRepository(conn).get_by_session_token(
|
||||
session_token
|
||||
)
|
||||
if not session:
|
||||
raise ValueError(
|
||||
f"Invalid session token ({session_token_fingerprint(session_token)})"
|
||||
)
|
||||
|
||||
token_info = session.get("token_info")
|
||||
if not token_info:
|
||||
raise ValueError("Session missing token information")
|
||||
|
||||
required_fields = ["access_token", "refresh_token"]
|
||||
missing_fields = [field for field in required_fields if field not in token_info or not token_info.get(field)]
|
||||
if missing_fields:
|
||||
raise ValueError(f"Missing required token fields: {missing_fields}")
|
||||
|
||||
if 'token_uri' not in token_info:
|
||||
token_info['token_uri'] = 'https://oauth2.googleapis.com/token'
|
||||
|
||||
return token_info
|
||||
|
||||
except Exception as e:
|
||||
raise ValueError(f"Failed to retrieve Google Drive token information: {str(e)}")
|
||||
|
||||
def validate_credentials(self, credentials: Credentials) -> bool:
|
||||
"""
|
||||
Validate Google Drive credentials by making a test API call.
|
||||
|
||||
@@ -48,11 +48,13 @@ class GoogleDriveLoader(BaseConnectorLoader):
|
||||
'application/vnd.google-apps.spreadsheet': 'application/vnd.openxmlformats-officedocument.spreadsheetml.sheet'
|
||||
}
|
||||
|
||||
def __init__(self, session_token: str):
|
||||
def __init__(self, session_token: Optional[str] = None, *, connection_id: Optional[str] = None):
|
||||
self.auth = GoogleDriveAuth()
|
||||
self.session_token = session_token
|
||||
|
||||
token_info = self.auth.get_token_info_from_session(session_token)
|
||||
_, token_info = self._load_token_info(session_token, connection_id)
|
||||
# Google refresh tokens do not rotate, so the credentials object may
|
||||
# renew the access token in memory for the rest of this run.
|
||||
self.credentials = self.auth.create_credentials_from_token_info(token_info)
|
||||
|
||||
try:
|
||||
|
||||
@@ -5,7 +5,6 @@ from typing import Optional, Dict, Any
|
||||
from msal import ConfidentialClientApplication
|
||||
|
||||
from docsgpt.core.settings import settings
|
||||
from docsgpt.parser.connectors._auth_utils import session_token_fingerprint
|
||||
from docsgpt.parser.connectors.base import BaseConnectorAuth
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -76,41 +75,6 @@ class SharePointAuth(BaseConnectorAuth):
|
||||
|
||||
return self.map_token_response(result)
|
||||
|
||||
def get_token_info_from_session(self, session_token: str) -> Dict[str, Any]:
|
||||
try:
|
||||
from docsgpt.storage.db.repositories.connector_sessions import (
|
||||
ConnectorSessionsRepository,
|
||||
)
|
||||
from docsgpt.storage.db.session import db_readonly
|
||||
|
||||
with db_readonly() as conn:
|
||||
session = ConnectorSessionsRepository(conn).get_by_session_token(
|
||||
session_token
|
||||
)
|
||||
|
||||
if not session:
|
||||
raise ValueError(
|
||||
f"Invalid session token ({session_token_fingerprint(session_token)})"
|
||||
)
|
||||
|
||||
token_info = session.get("token_info")
|
||||
if not token_info:
|
||||
raise ValueError("Session missing token information")
|
||||
|
||||
required_fields = ["access_token", "refresh_token"]
|
||||
missing_fields = [field for field in required_fields if field not in token_info or not token_info.get(field)]
|
||||
if missing_fields:
|
||||
raise ValueError(f"Missing required token fields: {missing_fields}")
|
||||
|
||||
if 'token_uri' not in token_info:
|
||||
token_info['token_uri'] = f"https://login.microsoftonline.com/{settings.MICROSOFT_TENANT_ID}/oauth2/v2.0/token"
|
||||
|
||||
return token_info
|
||||
|
||||
except Exception as e:
|
||||
logger.error("Failed to retrieve token from session: %s", e)
|
||||
raise ValueError(f"Failed to retrieve SharePoint token information: {str(e)}")
|
||||
|
||||
def is_token_expired(self, token_info: Dict[str, Any]) -> bool:
|
||||
if not token_info:
|
||||
return True
|
||||
|
||||
@@ -26,8 +26,7 @@ def _retry_on_auth_failure(func):
|
||||
if e.response is not None and e.response.status_code in (401, 403):
|
||||
logging.info(f"Auth failure in {func.__name__}, refreshing token and retrying")
|
||||
try:
|
||||
new_token_info = self.auth.refresh_access_token(self.refresh_token)
|
||||
self.access_token = new_token_info.get('access_token')
|
||||
self._apply_token_info(self._refresh_rejected_token(self.access_token))
|
||||
except Exception as refresh_error:
|
||||
raise ValueError(
|
||||
f"Authentication failed and could not be refreshed: {refresh_error}"
|
||||
@@ -63,13 +62,12 @@ class SharePointLoader(BaseConnectorLoader):
|
||||
|
||||
GRAPH_API_BASE = "https://graph.microsoft.com/v1.0"
|
||||
|
||||
def __init__(self, session_token: str):
|
||||
def __init__(self, session_token: Optional[str] = None, *, connection_id: Optional[str] = None):
|
||||
self.auth = SharePointAuth()
|
||||
self.session_token = session_token
|
||||
|
||||
token_info = self.auth.get_token_info_from_session(session_token)
|
||||
self.access_token = token_info.get('access_token')
|
||||
self.refresh_token = token_info.get('refresh_token')
|
||||
_, token_info = self._load_token_info(session_token, connection_id)
|
||||
self._apply_token_info(token_info)
|
||||
self.allows_shared_content = token_info.get('allows_shared_content', False)
|
||||
|
||||
if not self.access_token:
|
||||
@@ -77,6 +75,10 @@ class SharePointLoader(BaseConnectorLoader):
|
||||
|
||||
self.next_page_token = None
|
||||
|
||||
def _apply_token_info(self, token_info: Dict[str, Any]) -> None:
|
||||
self.access_token = token_info.get('access_token')
|
||||
self.expiry = token_info.get('expiry')
|
||||
|
||||
def _get_headers(self) -> Dict[str, str]:
|
||||
return {
|
||||
'Authorization': f'Bearer {self.access_token}',
|
||||
@@ -87,12 +89,14 @@ class SharePointLoader(BaseConnectorLoader):
|
||||
if not self.access_token:
|
||||
raise ValueError("No access token available")
|
||||
|
||||
token_info = {'access_token': self.access_token, 'expiry': None}
|
||||
if self.auth.is_token_expired(token_info):
|
||||
# The connection service refreshes under a row lock and stores the
|
||||
# rotated refresh token; refreshing here would spend it and lose it.
|
||||
if self.auth.is_token_expired({'access_token': self.access_token, 'expiry': self.expiry}):
|
||||
logging.info("Token expired, attempting refresh")
|
||||
try:
|
||||
new_token_info = self.auth.refresh_access_token(self.refresh_token)
|
||||
self.access_token = new_token_info.get('access_token')
|
||||
from docsgpt.connectors import service
|
||||
|
||||
self._apply_token_info(service.get_valid_token_info(self.connection_id))
|
||||
except Exception:
|
||||
raise ValueError("Failed to refresh access token")
|
||||
|
||||
|
||||
@@ -3,7 +3,7 @@ import logging
|
||||
import mimetypes
|
||||
import time
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from typing import Dict, List, Optional, Tuple
|
||||
from typing import Any, Dict, List, Optional, Tuple, Union
|
||||
|
||||
import requests
|
||||
|
||||
@@ -39,6 +39,14 @@ SKIP_SUFFIXES = (
|
||||
)
|
||||
|
||||
|
||||
class GitHubTokenRejected(PermissionError):
|
||||
"""GitHub answered 401 to a user's own token (revoked or expired)."""
|
||||
|
||||
|
||||
class PrivateRepositoryError(PermissionError):
|
||||
"""A repository the instance-wide token may not read for a user."""
|
||||
|
||||
|
||||
class GitHubLoader(BaseRemote):
|
||||
"""Load a GitHub repository's text files as ``Document`` objects.
|
||||
|
||||
@@ -48,15 +56,74 @@ class GitHubLoader(BaseRemote):
|
||||
downloads the survivors in parallel.
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
self.access_token = settings.GITHUB_ACCESS_TOKEN
|
||||
self.headers = {
|
||||
"Authorization": f"token {self.access_token}",
|
||||
"Accept": "application/vnd.github.v3+json"
|
||||
} if self.access_token else {
|
||||
"Accept": "application/vnd.github.v3+json"
|
||||
}
|
||||
return
|
||||
def __init__(self, access_token: Optional[str] = None):
|
||||
"""Create a loader.
|
||||
|
||||
Args:
|
||||
access_token: A user's own GitHub token (from their GitHub
|
||||
connection). Without one the instance-wide
|
||||
``GITHUB_ACCESS_TOKEN`` is used, and only for public
|
||||
repositories.
|
||||
"""
|
||||
self._user_token = access_token
|
||||
self._use_token(access_token or settings.GITHUB_ACCESS_TOKEN)
|
||||
|
||||
def _use_token(self, token: Optional[str]) -> None:
|
||||
"""Send ``token`` (or nothing) on every following request."""
|
||||
self.access_token = token
|
||||
self.headers = {"Accept": "application/vnd.github.v3+json"}
|
||||
if token:
|
||||
self.headers["Authorization"] = f"Bearer {token}"
|
||||
|
||||
@staticmethod
|
||||
def _parse_inputs(inputs: Union[str, Dict[str, Any]]) -> Tuple[str, Optional[str]]:
|
||||
"""``(repo_url, token)`` from a URL string or a connection's loader input.
|
||||
|
||||
A source synced from a GitHub connection gets a dict: the repository
|
||||
(``repo_url`` or ``url``) plus the connection's ``access_token``.
|
||||
"""
|
||||
if isinstance(inputs, dict):
|
||||
repo = inputs.get("repo_url") or inputs.get("url") or ""
|
||||
token = inputs.get("access_token") or inputs.get("token")
|
||||
return str(repo), (str(token) if token else None)
|
||||
return str(inputs or ""), None
|
||||
|
||||
def ensure_instance_token_allowed(self, repo_name: str) -> None:
|
||||
"""Refuse a repository the instance-wide token must not read.
|
||||
|
||||
``GITHUB_ACCESS_TOKEN`` belongs to the server, not to the user asking,
|
||||
and may see private repositories. Without a user's own token it is
|
||||
used only for public ones (for their higher rate limit); anything
|
||||
else needs the user's GitHub connection.
|
||||
|
||||
Raises:
|
||||
PrivateRepositoryError: The repository is private, internal, or
|
||||
not visible to the token.
|
||||
"""
|
||||
if self._user_token or not self.access_token:
|
||||
# A user's own token, or anonymous requests (which only ever see
|
||||
# public repositories).
|
||||
return
|
||||
url = f"https://api.github.com/repos/{repo_name}"
|
||||
response = requests.get(url, headers=self.headers, timeout=30)
|
||||
if response.status_code == 401:
|
||||
# A stale instance token: public repositories still read anonymously.
|
||||
logger.warning("GitHub rejected GITHUB_ACCESS_TOKEN (401); reading %s anonymously", repo_name)
|
||||
self._use_token(None)
|
||||
return
|
||||
if response.status_code == 200:
|
||||
try:
|
||||
data = response.json()
|
||||
except ValueError:
|
||||
data = {}
|
||||
if isinstance(data, dict) and data.get("private") is False and (
|
||||
data.get("visibility") or "public"
|
||||
) == "public":
|
||||
return
|
||||
raise PrivateRepositoryError(
|
||||
f"{repo_name} is private or could not be found. Connect your GitHub account "
|
||||
"under Connectors to sync private repositories."
|
||||
)
|
||||
|
||||
def is_text_file(self, file_path: str) -> bool:
|
||||
"""Determine if a file is a text file based on extension."""
|
||||
@@ -344,6 +411,10 @@ class GitHubLoader(BaseRemote):
|
||||
raise
|
||||
# If we can't parse the response, raise the original error
|
||||
response.raise_for_status()
|
||||
elif response.status_code == 401 and self._user_token:
|
||||
# The user's own token: a retry without it could only read
|
||||
# public repositories, which is not what they asked for.
|
||||
raise GitHubTokenRejected("GitHub rejected the connection's token. Reconnect GitHub to continue.")
|
||||
elif response.status_code == 401 and self.access_token:
|
||||
# An expired or revoked PAT makes even public repos 401, which
|
||||
# is strictly worse than not sending one. Retry unauthenticated
|
||||
@@ -404,14 +475,30 @@ class GitHubLoader(BaseRemote):
|
||||
paths = self.fetch_repo_files(repo_name)
|
||||
return self.select_files([(p, 0) for p in paths])
|
||||
|
||||
def load_data(self, repo_url: str) -> List[Document]:
|
||||
"""Load every ingestable text file in ``repo_url`` as a Document."""
|
||||
def load_data(self, inputs: Union[str, Dict[str, Any]]) -> List[Document]:
|
||||
"""Load every ingestable text file of a repository as a Document.
|
||||
|
||||
Args:
|
||||
inputs: The repository URL, or ``{"repo_url", "access_token"}``
|
||||
for a source synced from a GitHub connection.
|
||||
|
||||
Raises:
|
||||
ValueError: ``inputs`` names no github.com repository.
|
||||
PrivateRepositoryError: Only the instance-wide token is available
|
||||
and the repository is not public.
|
||||
GitHubTokenRejected: GitHub refused the connection's token.
|
||||
"""
|
||||
repo_url, token = self._parse_inputs(inputs)
|
||||
if token:
|
||||
self._user_token = token
|
||||
self._use_token(token)
|
||||
repo_name = self.normalize_repo(repo_url)
|
||||
if not repo_name or "/" not in repo_name:
|
||||
raise ValueError(
|
||||
f"Not a valid GitHub repository: {repo_url!r}. "
|
||||
"Expected a github.com URL like https://github.com/owner/name."
|
||||
)
|
||||
self.ensure_instance_token_allowed(repo_name)
|
||||
branch = self.get_default_branch(repo_name)
|
||||
files = self._list_candidate_files(repo_name, branch)
|
||||
logger.info(
|
||||
|
||||
@@ -0,0 +1,234 @@
|
||||
"""Load Linear issues and documents as Knowledge, through Linear's MCP server.
|
||||
|
||||
A Linear source names teams and projects (see
|
||||
``docsgpt.connectors.linear.normalize_selection``). Each issue becomes one
|
||||
document: its identifier and title, state, assignee, priority, labels,
|
||||
description and, optionally, its comments. With ``include_documents`` the
|
||||
picked projects' Linear documents come too. Every document carries its
|
||||
Linear URL as ``source``, so answers cite the issue.
|
||||
|
||||
Each sync reads everything again (up to ``MAX_ISSUES`` issues and
|
||||
``MAX_DOCUMENTS`` documents): the index of a synced source is rebuilt
|
||||
whole, so there is nothing to merge an incremental read into.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import re
|
||||
from typing import Any, Optional
|
||||
|
||||
from docsgpt.connectors import linear
|
||||
from docsgpt.connectors.mcp import MCPToolError, run_connection_session
|
||||
from docsgpt.parser.remote.base import BaseRemote
|
||||
from docsgpt.parser.schema.base import Document
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_UNSAFE_PATH = re.compile(r"[\\/\x00-\x1f]+")
|
||||
|
||||
|
||||
def _connection(connection_id: str) -> Optional[dict]:
|
||||
from docsgpt.storage.db.repositories.connector_sessions import ConnectorSessionsRepository
|
||||
from docsgpt.storage.db.session import db_readonly
|
||||
|
||||
with db_readonly() as conn:
|
||||
return ConnectorSessionsRepository(conn).get(str(connection_id))
|
||||
|
||||
|
||||
def _segment(value: str, fallback: str) -> str:
|
||||
"""A file tree segment: no slashes or control characters, never empty."""
|
||||
cleaned = _UNSAFE_PATH.sub(" ", value or "").strip().strip(".")
|
||||
return cleaned[:120] or fallback
|
||||
|
||||
|
||||
def _day(value: Any) -> str:
|
||||
return str(value)[:10] if value else ""
|
||||
|
||||
|
||||
class LinearLoader(BaseRemote):
|
||||
"""Read a Linear source's issues and documents with its connection's sign-in."""
|
||||
|
||||
def load_data(self, inputs: dict) -> list[Document]:
|
||||
"""Read what ``inputs`` picks, signed in with its ``connection_id``.
|
||||
|
||||
Args:
|
||||
inputs: The source's selection plus ``connection_id``, merged in
|
||||
by the worker from the source's connection.
|
||||
|
||||
Raises:
|
||||
ValueError: No connection, or nothing picked.
|
||||
docsgpt.connectors.service.ConnectionUnavailable: The connection
|
||||
needs reconnecting.
|
||||
"""
|
||||
inputs = inputs if isinstance(inputs, dict) else {}
|
||||
connection_id = inputs.get("connection_id")
|
||||
connection = _connection(connection_id) if connection_id else None
|
||||
if connection is None:
|
||||
raise ValueError("A Linear source needs its Linear connection")
|
||||
selection = linear.normalize_selection(inputs)
|
||||
return run_connection_session(
|
||||
connection, linear.mcp_url(), lambda session: self.collect(session, selection),
|
||||
)
|
||||
|
||||
async def collect(self, session: Any, selection: dict) -> list[Document]:
|
||||
"""The documents for ``selection``, read over an open MCP session."""
|
||||
issues: dict[str, dict] = {}
|
||||
filters = [(("team", "teamId"), team["id"]) for team in selection["teams"]]
|
||||
filters += [(("project", "projectId"), project["id"]) for project in selection["projects"]]
|
||||
for names, value in filters:
|
||||
remaining = linear.MAX_ISSUES - len(issues)
|
||||
if remaining <= 0:
|
||||
logger.info("Linear sync stopped at %d issues", linear.MAX_ISSUES)
|
||||
break
|
||||
async for issue in linear.paged(
|
||||
session, "list_issues", "issues", {names: value}, limit=remaining,
|
||||
extra={"orderBy": "updatedAt", "includeArchived": False},
|
||||
):
|
||||
key = linear.issue_identifier(issue) or str(issue.get("id") or "")
|
||||
if key:
|
||||
issues.setdefault(key, issue)
|
||||
documents = []
|
||||
for key, issue in issues.items():
|
||||
# One issue Linear will not read (deleted since it was listed, or
|
||||
# out of the token's reach) loses that detail, not the whole sync.
|
||||
if linear.is_truncated(issue.get("description")) or issue.get("descriptionTruncated"):
|
||||
try:
|
||||
issue = await self._full_issue(session, key, issue)
|
||||
except MCPToolError as exc:
|
||||
logger.warning("Linear sync keeps %s as listed: %s", key, exc)
|
||||
comments = []
|
||||
if selection["include_comments"]:
|
||||
try:
|
||||
comments = await self._comments(session, key)
|
||||
except MCPToolError as exc:
|
||||
logger.warning("Linear sync leaves out the comments of %s: %s", key, exc)
|
||||
documents.append(self._issue_document(issue, comments, selection))
|
||||
if selection["include_documents"]:
|
||||
documents.extend(await self._project_documents(session, selection["projects"]))
|
||||
return documents
|
||||
|
||||
async def _full_issue(self, session: Any, key: str, issue: dict) -> dict:
|
||||
schema = await session.input_schema("get_issue")
|
||||
if schema is None:
|
||||
return issue
|
||||
full = linear.unwrap(await session.call("get_issue", {linear.argument(schema, "id", "issueId"): key}), "issue")
|
||||
return {**issue, **{k: v for k, v in full.items() if v not in (None, "")}}
|
||||
|
||||
async def _comments(self, session: Any, key: str) -> list[dict]:
|
||||
"""An issue's comments, oldest first as Linear lists them; none when Linear cannot list them."""
|
||||
schema = await session.input_schema("list_comments")
|
||||
names = ("issueId", "id", "issue")
|
||||
if schema is None or linear.argument(schema, *names) is None:
|
||||
return []
|
||||
return [
|
||||
comment async for comment in linear.paged(
|
||||
session, "list_comments", "comments", {names: key}, limit=linear.MAX_COMMENTS,
|
||||
)
|
||||
]
|
||||
|
||||
def _issue_document(self, issue: dict, comments: list[dict], selection: dict) -> Document:
|
||||
identifier = linear.issue_identifier(issue) or str(issue.get("id") or "")
|
||||
title = str(issue.get("title") or "").strip() or identifier
|
||||
heading = f"{identifier}: {title}" if identifier else title
|
||||
facts = [
|
||||
("State", linear.name_of(issue.get("status") or issue.get("state"))),
|
||||
("Assignee", linear.name_of(issue.get("assignee"))),
|
||||
("Priority", linear.priority_of(issue.get("priority"))),
|
||||
("Labels", ", ".join(linear.names_of(issue.get("labels")))),
|
||||
("Project", linear.name_of(issue.get("project"))),
|
||||
("Team", linear.name_of(issue.get("team"))),
|
||||
("Due", _day(issue.get("dueDate"))),
|
||||
("Created", _day(issue.get("createdAt"))),
|
||||
("Updated", _day(issue.get("updatedAt"))),
|
||||
("Link", str(issue.get("url") or "")),
|
||||
]
|
||||
lines = [f"# {heading}", ""]
|
||||
lines += [f"- {label}: {value}" for label, value in facts if value]
|
||||
description = str(issue.get("description") or "").strip()
|
||||
if description:
|
||||
lines += ["", "## Description", "", description]
|
||||
written = [c for c in comments if str(c.get("body") or "").strip()]
|
||||
if written:
|
||||
lines += ["", "## Comments"]
|
||||
for comment in written:
|
||||
author = linear.name_of(comment.get("user") or comment.get("author")) or "Someone"
|
||||
when = _day(comment.get("createdAt"))
|
||||
lines += ["", f"**{author}**" + (f" · {when}" if when else ""), str(comment["body"]).strip()]
|
||||
return Document(
|
||||
text="\n".join(lines).strip() + "\n",
|
||||
extra_info={
|
||||
"title": heading,
|
||||
"source": str(issue.get("url") or ""),
|
||||
"file_path": f"{self._team_folder(issue, identifier, selection)}/{_segment(identifier, 'issue')}.md",
|
||||
"linear_id": str(issue.get("id") or identifier),
|
||||
"updated_at": str(issue.get("updatedAt") or ""),
|
||||
},
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _team_folder(issue: dict, identifier: str, selection: dict) -> str:
|
||||
"""The folder an issue is filed under: its team's key (``ENG``), else its team's name."""
|
||||
team = issue.get("team")
|
||||
key = team.get("key") if isinstance(team, dict) else None
|
||||
if not key and "-" in identifier:
|
||||
key = identifier.rsplit("-", 1)[0]
|
||||
return _segment(str(key or linear.name_of(team) or ""), "Issues")
|
||||
|
||||
async def _project_documents(self, session: Any, projects: list[dict]) -> list[Document]:
|
||||
"""The Linear documents of ``projects``, each read in full, up to ``MAX_DOCUMENTS``."""
|
||||
documents: list[Document] = []
|
||||
paths: set[str] = set()
|
||||
get_schema = await session.input_schema("get_document")
|
||||
for project in projects:
|
||||
remaining = linear.MAX_DOCUMENTS - len(documents)
|
||||
if remaining <= 0:
|
||||
break
|
||||
async for record in linear.paged(
|
||||
session, "list_documents", "documents", {("projectId", "project"): project["id"]}, limit=remaining,
|
||||
):
|
||||
content = record.get("content")
|
||||
if (not content or linear.is_truncated(content)) and get_schema is not None and record.get("id"):
|
||||
try:
|
||||
full = linear.unwrap(
|
||||
await session.call("get_document", {linear.argument(get_schema, "id", "documentId"):
|
||||
str(record["id"])}),
|
||||
"document",
|
||||
)
|
||||
except MCPToolError as exc:
|
||||
# A document that cannot be read is skipped, not the sync.
|
||||
logger.warning("Linear sync skips document %s: %s", record.get("id"), exc)
|
||||
continue
|
||||
record = {**record, **{k: v for k, v in full.items() if v not in (None, "")}}
|
||||
document = self._linear_document(record, project)
|
||||
path = document.extra_info["file_path"]
|
||||
if path in paths:
|
||||
# Two documents with one title stay two files.
|
||||
document.extra_info["file_path"] = f"{path[:-3]} ({str(record.get('id') or len(paths))[:8]}).md"
|
||||
paths.add(document.extra_info["file_path"])
|
||||
documents.append(document)
|
||||
return documents
|
||||
|
||||
@staticmethod
|
||||
def _linear_document(record: dict, project: dict) -> Document:
|
||||
title = str(record.get("title") or "").strip() or "Untitled document"
|
||||
content = str(record.get("content") or "").strip()
|
||||
project_name = project.get("name") or linear.name_of(record.get("project"))
|
||||
lines = [f"# {title}", ""]
|
||||
if project_name:
|
||||
lines.append(f"- Project: {project_name}")
|
||||
if record.get("url"):
|
||||
lines.append(f"- Link: {record['url']}")
|
||||
if content:
|
||||
lines += ["", content]
|
||||
folder = _segment(project_name or "", "Documents")
|
||||
return Document(
|
||||
text="\n".join(lines).strip() + "\n",
|
||||
extra_info={
|
||||
"title": title,
|
||||
"source": str(record.get("url") or ""),
|
||||
"file_path": f"{folder}/Documents/{_segment(title, 'Document')}.md",
|
||||
"linear_id": str(record.get("id") or ""),
|
||||
"updated_at": str(record.get("updatedAt") or ""),
|
||||
},
|
||||
)
|
||||
@@ -5,6 +5,7 @@ from docsgpt.parser.remote.crawler_loader import CrawlerLoader
|
||||
from docsgpt.parser.remote.web_loader import WebLoader
|
||||
from docsgpt.parser.remote.reddit_loader import RedditPostsLoaderRemote
|
||||
from docsgpt.parser.remote.github_loader import GitHubLoader
|
||||
from docsgpt.parser.remote.linear_loader import LinearLoader
|
||||
from docsgpt.parser.remote.s3_loader import S3Loader
|
||||
|
||||
|
||||
@@ -26,6 +27,7 @@ class RemoteCreator:
|
||||
"reddit": RedditPostsLoaderRemote,
|
||||
"github": GitHubLoader,
|
||||
"s3": S3Loader,
|
||||
"linear": LinearLoader,
|
||||
}
|
||||
|
||||
@classmethod
|
||||
@@ -88,5 +90,5 @@ def normalize_remote_data(source_type, remote_data):
|
||||
return json.dumps(remote_data)
|
||||
return remote_data
|
||||
|
||||
# s3's loader accepts a dict or JSON string; pass it through unchanged.
|
||||
# s3's and linear's loaders accept a dict or JSON string; pass it through unchanged.
|
||||
return remote_data
|
||||
@@ -2,6 +2,17 @@ from abc import ABC, abstractmethod
|
||||
|
||||
|
||||
class BaseRetriever(ABC):
|
||||
@property
|
||||
def _connector_labels(self):
|
||||
"""Per-retriever memo of each source's connector ("From Google Drive")."""
|
||||
cache = self.__dict__.get("_connector_label_cache")
|
||||
if cache is None:
|
||||
from docsgpt.connectors.attribution import ConnectorLabelCache
|
||||
|
||||
cache = ConnectorLabelCache()
|
||||
self.__dict__["_connector_label_cache"] = cache
|
||||
return cache
|
||||
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
|
||||
@@ -385,7 +385,11 @@ class ClassicRAG(BaseRetriever):
|
||||
doc_tokens = num_tokens_from_string(doc_text_with_header)
|
||||
|
||||
if cumulative_tokens + doc_tokens < token_budget:
|
||||
entry = {"text": page_content, **labels}
|
||||
entry = {
|
||||
"text": page_content,
|
||||
**labels,
|
||||
**self._connector_labels.for_source(vectorstore_id),
|
||||
}
|
||||
if self.include_scores:
|
||||
entry["score"] = score
|
||||
entry["score_kind"] = score_kind
|
||||
|
||||
@@ -584,7 +584,7 @@ class GraphRAGRetriever(BaseRetriever):
|
||||
doc_tokens = num_tokens_from_string(f"{labels['filename']}\n{text}")
|
||||
if cumulative_tokens + doc_tokens >= token_budget:
|
||||
break
|
||||
docs.append({"text": text, **labels})
|
||||
docs.append({"text": text, **labels, **self._connector_labels.for_source(source_id)})
|
||||
cumulative_tokens += doc_tokens
|
||||
return docs
|
||||
|
||||
|
||||
@@ -1,11 +1,17 @@
|
||||
import base64
|
||||
import functools
|
||||
import hashlib
|
||||
import hmac
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
from typing import Optional
|
||||
|
||||
from cryptography.hazmat.backends import default_backend
|
||||
from cryptography.hazmat.primitives import hashes
|
||||
from cryptography.hazmat.primitives.ciphers import algorithms, Cipher, modes
|
||||
from cryptography.hazmat.primitives.ciphers.aead import AESGCM
|
||||
from cryptography.hazmat.primitives.kdf.hkdf import HKDF
|
||||
from cryptography.hazmat.primitives.kdf.pbkdf2 import PBKDF2HMAC
|
||||
|
||||
from docsgpt.core.settings import settings
|
||||
@@ -86,3 +92,134 @@ def _pad_data(data: bytes) -> bytes:
|
||||
def _unpad_data(data: bytes) -> bytes:
|
||||
padding_len = data[-1]
|
||||
return data[:-padding_len]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Envelope v2: connection credentials
|
||||
# ---------------------------------------------------------------------------
|
||||
#
|
||||
# ``v2:<key_id>:<base64(salt | nonce | ciphertext+tag)>``
|
||||
#
|
||||
# AES-256-GCM, so a tampered blob fails to decrypt instead of returning
|
||||
# garbage. The master key is derived once per process from
|
||||
# ENCRYPTION_SECRET_KEY (PBKDF2, cached); each record gets its own key from
|
||||
# HKDF(master, salt, owner id), which keeps the v1 owner binding without
|
||||
# paying 100k PBKDF2 iterations on every token read in the worker. The owner
|
||||
# id is also the GCM associated data, so a blob copied onto another user's
|
||||
# row does not decrypt. ``key_id`` names the master key, so a blob written
|
||||
# under ENCRYPTION_SECRET_KEY_PREVIOUS is still readable during a rotation.
|
||||
|
||||
_V2_PREFIX = "v2"
|
||||
_V2_MASTER_SALT = b"docsgpt-credentials-v2"
|
||||
_V2_ITERATIONS = 200_000
|
||||
_V2_SALT_BYTES = 16
|
||||
_V2_NONCE_BYTES = 12
|
||||
DEFAULT_ENCRYPTION_KEY = "default-docsgpt-encryption-key"
|
||||
|
||||
|
||||
class CredentialDecryptionError(Exception):
|
||||
"""A stored credential could not be decrypted (wrong key, tampering, bad format)."""
|
||||
|
||||
|
||||
@functools.lru_cache(maxsize=8)
|
||||
def _master_key(secret: str) -> bytes:
|
||||
kdf = PBKDF2HMAC(
|
||||
algorithm=hashes.SHA256(),
|
||||
length=32,
|
||||
salt=_V2_MASTER_SALT,
|
||||
iterations=_V2_ITERATIONS,
|
||||
backend=default_backend(),
|
||||
)
|
||||
return kdf.derive(secret.encode())
|
||||
|
||||
|
||||
def _key_id(master: bytes) -> str:
|
||||
return hmac.new(master, b"docsgpt-key-id", hashlib.sha256).hexdigest()[:8]
|
||||
|
||||
|
||||
def _record_key(master: bytes, owner_id: str, salt: bytes) -> bytes:
|
||||
return HKDF(
|
||||
algorithm=hashes.SHA256(),
|
||||
length=32,
|
||||
salt=salt,
|
||||
info=b"docsgpt-v2|" + owner_id.encode(),
|
||||
backend=default_backend(),
|
||||
).derive(master)
|
||||
|
||||
|
||||
def _candidate_keys() -> dict[str, bytes]:
|
||||
"""Master keys this process can decrypt with, by key id (current first)."""
|
||||
keys: dict[str, bytes] = {}
|
||||
for secret in (settings.ENCRYPTION_SECRET_KEY, settings.ENCRYPTION_SECRET_KEY_PREVIOUS):
|
||||
if secret:
|
||||
master = _master_key(secret)
|
||||
keys.setdefault(_key_id(master), master)
|
||||
return keys
|
||||
|
||||
|
||||
def current_key_id() -> str:
|
||||
"""Key id of ENCRYPTION_SECRET_KEY, as written into new v2 blobs."""
|
||||
return _key_id(_master_key(settings.ENCRYPTION_SECRET_KEY))
|
||||
|
||||
|
||||
def is_default_encryption_key() -> bool:
|
||||
"""Whether ENCRYPTION_SECRET_KEY is still the public default."""
|
||||
return settings.ENCRYPTION_SECRET_KEY == DEFAULT_ENCRYPTION_KEY
|
||||
|
||||
|
||||
def encrypt_json(data: dict, owner_id: str) -> str:
|
||||
"""Encrypt ``data`` for ``owner_id`` into a v2 envelope.
|
||||
|
||||
Args:
|
||||
data: JSON-serialisable credentials.
|
||||
owner_id: The user the credentials belong to; decryption needs it.
|
||||
|
||||
Returns:
|
||||
The ``v2:<key_id>:<payload>`` string.
|
||||
"""
|
||||
master = _master_key(settings.ENCRYPTION_SECRET_KEY)
|
||||
key_id = _key_id(master)
|
||||
salt = os.urandom(_V2_SALT_BYTES)
|
||||
nonce = os.urandom(_V2_NONCE_BYTES)
|
||||
key = _record_key(master, owner_id, salt)
|
||||
plaintext = json.dumps(data, separators=(",", ":")).encode()
|
||||
ciphertext = AESGCM(key).encrypt(nonce, plaintext, owner_id.encode())
|
||||
payload = base64.b64encode(salt + nonce + ciphertext).decode()
|
||||
return f"{_V2_PREFIX}:{key_id}:{payload}"
|
||||
|
||||
|
||||
def envelope_key_id(blob: str) -> Optional[str]:
|
||||
"""The key id a v2 blob was written with, or None for anything else."""
|
||||
parts = (blob or "").split(":", 2)
|
||||
if len(parts) != 3 or parts[0] != _V2_PREFIX:
|
||||
return None
|
||||
return parts[1]
|
||||
|
||||
|
||||
def decrypt_json(blob: str, owner_id: str) -> dict:
|
||||
"""Decrypt a v2 envelope written for ``owner_id``.
|
||||
|
||||
Raises:
|
||||
CredentialDecryptionError: The blob is malformed, was written with a
|
||||
key this process does not have, belongs to another owner, or was
|
||||
tampered with.
|
||||
"""
|
||||
key_id = envelope_key_id(blob)
|
||||
if key_id is None:
|
||||
raise CredentialDecryptionError("Not a v2 credential envelope")
|
||||
master = _candidate_keys().get(key_id)
|
||||
if master is None:
|
||||
raise CredentialDecryptionError("Credential was encrypted with an unknown key")
|
||||
try:
|
||||
raw = base64.b64decode(blob.split(":", 2)[2].encode(), validate=True)
|
||||
salt = raw[:_V2_SALT_BYTES]
|
||||
nonce = raw[_V2_SALT_BYTES:_V2_SALT_BYTES + _V2_NONCE_BYTES]
|
||||
ciphertext = raw[_V2_SALT_BYTES + _V2_NONCE_BYTES:]
|
||||
key = _record_key(master, owner_id, salt)
|
||||
plaintext = AESGCM(key).decrypt(nonce, ciphertext, owner_id.encode())
|
||||
data = json.loads(plaintext.decode())
|
||||
except Exception as exc:
|
||||
raise CredentialDecryptionError("Credential could not be decrypted") from exc
|
||||
if not isinstance(data, dict):
|
||||
raise CredentialDecryptionError("Credential payload is not an object")
|
||||
return data
|
||||
@@ -65,7 +65,8 @@ def _authorized_source_ids(conn, agent: Dict[str, Any], source_ids: List[str]) -
|
||||
source_ids: Ids extracted from that row.
|
||||
|
||||
A source the owner can't read still searches while the editor who
|
||||
attached it (its sponsor) can edit the agent and read the source.
|
||||
attached it (its sponsor) can edit the agent and still owns or edits
|
||||
the source.
|
||||
|
||||
Returns:
|
||||
list: The subset the agent's owner (or a live sponsor) may read.
|
||||
|
||||
@@ -233,6 +233,11 @@ user_tools_table = Table(
|
||||
Column("created_at", DateTime(timezone=True), nullable=False, server_default=func.now()),
|
||||
Column("updated_at", DateTime(timezone=True), nullable=False, server_default=func.now()),
|
||||
Column("legacy_mongo_id", Text),
|
||||
Column(
|
||||
"connection_id", UUID(as_uuid=True), ForeignKey("connector_sessions.id", ondelete="SET NULL"),
|
||||
),
|
||||
# Whose account a shared resource runs with: the owner's, or each member's.
|
||||
Column("credential_mode", Text, nullable=False, server_default="owner"),
|
||||
)
|
||||
|
||||
# A grantee's personal "In my chats" switch for a tool shared with them
|
||||
@@ -371,6 +376,13 @@ sources_table = Table(
|
||||
Column("created_at", DateTime(timezone=True), nullable=False, server_default=func.now()),
|
||||
Column("updated_at", DateTime(timezone=True), nullable=False, server_default=func.now()),
|
||||
Column("legacy_mongo_id", Text),
|
||||
Column(
|
||||
"connection_id", UUID(as_uuid=True), ForeignKey("connector_sessions.id", ondelete="SET NULL"),
|
||||
),
|
||||
# Whose account a shared resource runs with: the owner's, or each member's.
|
||||
Column("credential_mode", Text, nullable=False, server_default="owner"),
|
||||
# A wiki's owner lets API-key and widget runs edit it (off: read only).
|
||||
Column("wiki_outside_edits", Boolean, nullable=False, server_default="false"),
|
||||
)
|
||||
|
||||
agents_table = Table(
|
||||
@@ -621,6 +633,34 @@ connector_sessions_table = Table(
|
||||
Column("expires_at", DateTime(timezone=True)),
|
||||
Column("created_at", DateTime(timezone=True), nullable=False, server_default=func.now()),
|
||||
Column("legacy_mongo_id", Text),
|
||||
# Added in ``0040_connections``: each row is a connection (one signed-in
|
||||
# account, one MCP server or one set of API credentials).
|
||||
Column("connector_key", Text),
|
||||
Column("display_name", Text),
|
||||
Column("account_label", Text),
|
||||
Column("auth_kind", Text),
|
||||
Column("updated_at", DateTime(timezone=True), nullable=False, server_default=func.now()),
|
||||
# Every secret of the connection, as one owner-bound v2 envelope.
|
||||
Column("encrypted_credentials", Text),
|
||||
Column("has_refresh_token", Boolean, nullable=False, server_default="false"),
|
||||
Column("scopes", JSONB, nullable=False, server_default="[]"),
|
||||
Column("last_error", Text),
|
||||
Column("last_used_at", DateTime(timezone=True)),
|
||||
# Added in ``0041_connection_account_name``: what the user calls the
|
||||
# account; ``account_label`` stays its identity.
|
||||
Column("account_name", Text),
|
||||
)
|
||||
|
||||
|
||||
connector_policies_table = Table(
|
||||
"connector_policies",
|
||||
metadata,
|
||||
Column("connector_key", Text, primary_key=True),
|
||||
# NULL: on when the connector has its server settings (see connectors.service).
|
||||
Column("enabled", Boolean),
|
||||
Column("credential_mode", Text, nullable=False, server_default="choose"),
|
||||
Column("updated_by", Text),
|
||||
Column("updated_at", DateTime(timezone=True), nullable=False, server_default=func.now()),
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -33,13 +33,17 @@ _SECRET_SUBSTRINGS = (
|
||||
"credential",
|
||||
"authorization",
|
||||
"bearer",
|
||||
# Connection secrets: an OAuth token_info blob, MCP token and client
|
||||
# registration dicts (``client_secret`` is covered by ``secret``).
|
||||
"token_info",
|
||||
"client_info",
|
||||
)
|
||||
|
||||
|
||||
def is_secret_key(key: str) -> bool:
|
||||
"""True when ``key`` names a credential that must not be persisted/returned."""
|
||||
k = key.lower()
|
||||
if k == "token":
|
||||
if k in ("token", "tokens"):
|
||||
return True
|
||||
return any(s in k for s in _SECRET_SUBSTRINGS)
|
||||
|
||||
|
||||
@@ -0,0 +1,72 @@
|
||||
"""Repository for ``connector_policies``: the admin's per-connector switches."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Optional
|
||||
|
||||
from sqlalchemy import Connection, text
|
||||
|
||||
from docsgpt.storage.db.base_repository import row_to_dict
|
||||
|
||||
CREDENTIAL_POLICIES = ("choose", "owner", "member")
|
||||
ALLOW_CUSTOM_MCP_KEY = "connectors.allow_custom_mcp"
|
||||
|
||||
|
||||
def allow_writes_key(connector_key: str) -> str:
|
||||
"""``app_metadata`` key of the admin's "let agents make changes" switch for a connector.
|
||||
|
||||
Only connectors whose tools opt into writes (GitHub) have one. It lives in
|
||||
``app_metadata`` like the custom MCP switch: ``"false"`` forbids writes,
|
||||
anything else (or nothing) allows them.
|
||||
"""
|
||||
return f"connectors.{connector_key}.allow_writes"
|
||||
|
||||
|
||||
class ConnectorPoliciesRepository:
|
||||
"""Whether a connector is enabled and which credential mode it forces."""
|
||||
|
||||
def __init__(self, conn: Connection) -> None:
|
||||
self._conn = conn
|
||||
|
||||
def all(self) -> dict[str, dict]:
|
||||
"""Every stored policy, by connector key. Missing keys use the defaults."""
|
||||
result = self._conn.execute(text("SELECT * FROM connector_policies"))
|
||||
return {row["connector_key"]: row for row in (row_to_dict(r) for r in result.fetchall())}
|
||||
|
||||
def get(self, connector_key: str) -> Optional[dict]:
|
||||
row = self._conn.execute(
|
||||
text("SELECT * FROM connector_policies WHERE connector_key = :key"), {"key": connector_key},
|
||||
).fetchone()
|
||||
return row_to_dict(row) if row is not None else None
|
||||
|
||||
def upsert(
|
||||
self,
|
||||
connector_key: str,
|
||||
*,
|
||||
enabled: Optional[bool] = None,
|
||||
credential_mode: Optional[str] = None,
|
||||
updated_by: Optional[str] = None,
|
||||
) -> dict:
|
||||
"""Set one connector's policy; fields left None keep their value.
|
||||
|
||||
A new row leaves ``enabled`` NULL unless it is given, so changing only
|
||||
the credential mode never switches a connector on.
|
||||
"""
|
||||
if credential_mode is not None and credential_mode not in CREDENTIAL_POLICIES:
|
||||
raise ValueError(f"unknown credential mode: {credential_mode!r}")
|
||||
row = self._conn.execute(
|
||||
text(
|
||||
"""
|
||||
INSERT INTO connector_policies (connector_key, enabled, credential_mode, updated_by)
|
||||
VALUES (:key, :enabled, COALESCE(:mode, 'choose'), :by)
|
||||
ON CONFLICT (connector_key) DO UPDATE SET
|
||||
enabled = COALESCE(:enabled, connector_policies.enabled),
|
||||
credential_mode = COALESCE(:mode, connector_policies.credential_mode),
|
||||
updated_by = :by,
|
||||
updated_at = now()
|
||||
RETURNING *
|
||||
"""
|
||||
),
|
||||
{"key": connector_key, "enabled": enabled, "mode": credential_mode, "by": updated_by},
|
||||
).fetchone()
|
||||
return row_to_dict(row)
|
||||
@@ -9,8 +9,13 @@ Shape notes:
|
||||
unique constraint on ``session_token``.
|
||||
* MCP sessions key off ``server_url`` instead — a single user may have
|
||||
multiple MCP servers, one row each. The composite unique index
|
||||
``(user_id, COALESCE(server_url, ''), provider)`` makes both patterns
|
||||
coexist without collision.
|
||||
``(user_id, provider, COALESCE(server_url, ''), COALESCE(account_label, ''))``
|
||||
makes both patterns coexist, and lets one user connect several accounts
|
||||
of the same provider (each row carries its own ``account_label``).
|
||||
* Every secret lives in ``encrypted_credentials``, written and read only by
|
||||
``docsgpt.connectors.service``; ``token_info`` and the ``tokens`` /
|
||||
``client_info`` keys of ``session_data`` are legacy plaintext that
|
||||
migration 0040 moved into it.
|
||||
* ``session_data`` remains a catch-all JSONB for driver-specific state
|
||||
(tokens that don't fit anywhere else, per-provider scratch data).
|
||||
Promoted columns (``session_token``, ``user_email``, ``status``,
|
||||
@@ -24,14 +29,16 @@ from typing import Any, Optional
|
||||
|
||||
from sqlalchemy import Connection, text
|
||||
|
||||
from docsgpt.storage.db.base_repository import row_to_dict
|
||||
from docsgpt.storage.db.base_repository import looks_like_uuid, row_to_dict
|
||||
from docsgpt.storage.db.serialization import PGNativeJSONEncoder
|
||||
|
||||
|
||||
_UPDATABLE_SCALARS = {
|
||||
"server_url", "session_token", "user_email", "status", "expires_at",
|
||||
"connector_key", "display_name", "account_label", "auth_kind",
|
||||
"encrypted_credentials", "has_refresh_token", "last_error", "last_used_at", "account_name",
|
||||
}
|
||||
_UPDATABLE_JSONB = {"session_data", "token_info"}
|
||||
_UPDATABLE_JSONB = {"session_data", "token_info", "scopes"}
|
||||
|
||||
|
||||
def _jsonb(value: Any) -> Any:
|
||||
@@ -70,7 +77,8 @@ class ConnectorSessionsRepository:
|
||||
) -> dict:
|
||||
"""Insert or update a connector session row.
|
||||
|
||||
Conflict key is ``(user_id, COALESCE(server_url, ''), provider)``
|
||||
Conflict key is the account index
|
||||
``(user_id, provider, COALESCE(server_url, ''), COALESCE(account_label, ''))``
|
||||
so MCP rows (per-server) and OAuth rows (per-provider) both get
|
||||
idempotent upsert semantics.
|
||||
"""
|
||||
@@ -86,7 +94,7 @@ class ConnectorSessionsRepository:
|
||||
:status, CAST(:token_info AS jsonb),
|
||||
CAST(:session_data AS jsonb), :expires_at, :legacy_mongo_id
|
||||
)
|
||||
ON CONFLICT (user_id, COALESCE(server_url, ''), provider)
|
||||
ON CONFLICT (user_id, provider, COALESCE(server_url, ''), COALESCE(account_label, ''))
|
||||
DO UPDATE SET
|
||||
session_token = COALESCE(EXCLUDED.session_token, connector_sessions.session_token),
|
||||
user_email = COALESCE(EXCLUDED.user_email, connector_sessions.user_email),
|
||||
@@ -176,11 +184,154 @@ class ConnectorSessionsRepository:
|
||||
|
||||
def list_for_user(self, user_id: str) -> list[dict]:
|
||||
result = self._conn.execute(
|
||||
text("SELECT * FROM connector_sessions WHERE user_id = :user_id"),
|
||||
text("SELECT * FROM connector_sessions WHERE user_id = :user_id ORDER BY created_at"),
|
||||
{"user_id": user_id},
|
||||
)
|
||||
return [row_to_dict(r) for r in result.fetchall()]
|
||||
|
||||
def create(
|
||||
self,
|
||||
user_id: str,
|
||||
provider: str,
|
||||
*,
|
||||
connector_key: str,
|
||||
auth_kind: str,
|
||||
display_name: Optional[str] = None,
|
||||
account_label: Optional[str] = None,
|
||||
server_url: Optional[str] = None,
|
||||
status: str = "connected",
|
||||
encrypted_credentials: Optional[str] = None,
|
||||
has_refresh_token: bool = False,
|
||||
) -> Optional[dict]:
|
||||
"""Insert a connection; return None when that account already exists."""
|
||||
result = self._conn.execute(
|
||||
text(
|
||||
"""
|
||||
INSERT INTO connector_sessions (
|
||||
user_id, provider, server_url, connector_key, auth_kind, display_name,
|
||||
account_label, status, encrypted_credentials, has_refresh_token, session_data
|
||||
)
|
||||
VALUES (
|
||||
:user_id, :provider, :server_url, :connector_key, :auth_kind, :display_name,
|
||||
:account_label, :status, :encrypted_credentials, :has_refresh_token, '{}'::jsonb
|
||||
)
|
||||
ON CONFLICT (user_id, provider, COALESCE(server_url, ''), COALESCE(account_label, ''))
|
||||
DO NOTHING
|
||||
RETURNING *
|
||||
"""
|
||||
),
|
||||
{
|
||||
"user_id": user_id,
|
||||
"provider": provider,
|
||||
"server_url": server_url,
|
||||
"connector_key": connector_key,
|
||||
"auth_kind": auth_kind,
|
||||
"display_name": display_name,
|
||||
"account_label": account_label,
|
||||
"status": status,
|
||||
"encrypted_credentials": encrypted_credentials,
|
||||
"has_refresh_token": has_refresh_token,
|
||||
},
|
||||
)
|
||||
row = result.fetchone()
|
||||
return row_to_dict(row) if row is not None else None
|
||||
|
||||
def find_account(
|
||||
self, user_id: str, provider: str, *, server_url: Optional[str], account_label: Optional[str],
|
||||
) -> Optional[dict]:
|
||||
"""The connection for one account, matching the account unique index."""
|
||||
result = self._conn.execute(
|
||||
text(
|
||||
"SELECT * FROM connector_sessions WHERE user_id = :user_id AND provider = :provider "
|
||||
"AND COALESCE(server_url, '') = COALESCE(:server_url, '') "
|
||||
"AND COALESCE(account_label, '') = COALESCE(:account_label, '')"
|
||||
),
|
||||
{"user_id": user_id, "provider": provider, "server_url": server_url, "account_label": account_label},
|
||||
)
|
||||
row = result.fetchone()
|
||||
return row_to_dict(row) if row is not None else None
|
||||
|
||||
def delete_by_id(self, connection_id: str) -> bool:
|
||||
"""Delete a connection row. Linked sources and tools keep existing (SET NULL)."""
|
||||
if not looks_like_uuid(connection_id):
|
||||
return False
|
||||
result = self._conn.execute(
|
||||
text("DELETE FROM connector_sessions WHERE id = CAST(:id AS uuid)"), {"id": str(connection_id)},
|
||||
)
|
||||
return result.rowcount > 0
|
||||
|
||||
def get(self, connection_id: str) -> Optional[dict]:
|
||||
"""Fetch a connection by id, whoever owns it. Callers authorise."""
|
||||
if not looks_like_uuid(connection_id):
|
||||
return None
|
||||
result = self._conn.execute(
|
||||
text("SELECT * FROM connector_sessions WHERE id = CAST(:id AS uuid)"),
|
||||
{"id": str(connection_id)},
|
||||
)
|
||||
row = result.fetchone()
|
||||
return row_to_dict(row) if row is not None else None
|
||||
|
||||
def get_for_user(self, connection_id: str, user_id: str) -> Optional[dict]:
|
||||
"""Fetch a connection only when ``user_id`` owns it."""
|
||||
row = self.get(connection_id)
|
||||
if row is None or row.get("user_id") != user_id:
|
||||
return None
|
||||
return row
|
||||
|
||||
def get_for_update(self, connection_id: str) -> Optional[dict]:
|
||||
"""Fetch and row-lock a connection until the transaction ends.
|
||||
|
||||
Token refresh holds this lock so two workers refreshing a rotating
|
||||
refresh token (Microsoft, Atlassian) cannot both spend it.
|
||||
"""
|
||||
if not looks_like_uuid(connection_id):
|
||||
return None
|
||||
result = self._conn.execute(
|
||||
text("SELECT * FROM connector_sessions WHERE id = CAST(:id AS uuid) FOR UPDATE"),
|
||||
{"id": str(connection_id)},
|
||||
)
|
||||
row = result.fetchone()
|
||||
return row_to_dict(row) if row is not None else None
|
||||
|
||||
def resource_counts(self, connection_ids: list[str]) -> dict[str, dict[str, int]]:
|
||||
"""Number of sources and tools linked to each connection id."""
|
||||
ids = [str(i) for i in connection_ids if looks_like_uuid(str(i))]
|
||||
counts: dict[str, dict[str, int]] = {i: {"sources": 0, "tools": 0} for i in ids}
|
||||
if not ids:
|
||||
return counts
|
||||
for table, key in (("sources", "sources"), ("user_tools", "tools")):
|
||||
result = self._conn.execute(
|
||||
text(
|
||||
f"SELECT connection_id, count(*) FROM {table} "
|
||||
"WHERE connection_id = ANY(CAST(:ids AS uuid[])) GROUP BY connection_id"
|
||||
),
|
||||
{"ids": ids},
|
||||
)
|
||||
for connection_id, count in result.fetchall():
|
||||
counts[str(connection_id)][key] = int(count)
|
||||
return counts
|
||||
|
||||
def list_sources(self, connection_id: str) -> list[dict]:
|
||||
"""Sources synced from a connection, newest first."""
|
||||
result = self._conn.execute(
|
||||
text(
|
||||
"SELECT id, name, type, date, sync_frequency, metadata, remote_data, user_id, file_path "
|
||||
"FROM sources WHERE connection_id = CAST(:id AS uuid) ORDER BY date DESC"
|
||||
),
|
||||
{"id": str(connection_id)},
|
||||
)
|
||||
return [row_to_dict(r) for r in result.fetchall()]
|
||||
|
||||
def list_tools(self, connection_id: str) -> list[dict]:
|
||||
"""Tools a connection provides, oldest first."""
|
||||
result = self._conn.execute(
|
||||
text(
|
||||
"SELECT * FROM user_tools WHERE connection_id = CAST(:id AS uuid) ORDER BY created_at"
|
||||
),
|
||||
{"id": str(connection_id)},
|
||||
)
|
||||
return [row_to_dict(r) for r in result.fetchall()]
|
||||
|
||||
def update(self, session_id: str, fields: dict) -> bool:
|
||||
"""Partial update by PG UUID."""
|
||||
filtered = {
|
||||
@@ -198,6 +349,7 @@ class ConnectorSessionsRepository:
|
||||
else:
|
||||
set_clauses.append(f"{col} = :{col}")
|
||||
params[col] = val
|
||||
set_clauses.append("updated_at = now()")
|
||||
result = self._conn.execute(
|
||||
text(
|
||||
f"UPDATE connector_sessions SET {', '.join(set_clauses)} "
|
||||
@@ -263,7 +415,7 @@ class ConnectorSessionsRepository:
|
||||
|
||||
Notes:
|
||||
The conflict target matches the table's composite unique
|
||||
constraint ``(user_id, COALESCE(server_url, ''), provider)``
|
||||
index ``(user_id, provider, COALESCE(server_url, ''), COALESCE(account_label, ''))``
|
||||
so MCP's per-URL rows and OAuth's single-row-per-user rows
|
||||
both upsert idempotently.
|
||||
"""
|
||||
@@ -284,7 +436,7 @@ class ConnectorSessionsRepository:
|
||||
:user_id, :provider, :server_url,
|
||||
CAST(:patch AS jsonb)
|
||||
)
|
||||
ON CONFLICT (user_id, COALESCE(server_url, ''), provider)
|
||||
ON CONFLICT (user_id, provider, COALESCE(server_url, ''), COALESCE(account_label, ''))
|
||||
DO UPDATE SET
|
||||
server_url = COALESCE(EXCLUDED.server_url, connector_sessions.server_url),
|
||||
session_data =
|
||||
|
||||
@@ -399,6 +399,30 @@ class SourcesRepository:
|
||||
)
|
||||
self._conn.execute(stmt)
|
||||
|
||||
def set_wiki_outside_edits(self, source_id: str, user_id: str, allowed: bool) -> bool:
|
||||
"""Record whether API-key and widget runs may edit this wiki.
|
||||
|
||||
Kept out of :meth:`update`'s columns so no route that forwards a
|
||||
request body can change it; only the owner-checked wiki settings
|
||||
route calls this.
|
||||
|
||||
Args:
|
||||
source_id: The source's UUID.
|
||||
user_id: The owner's id; the row is scoped to it.
|
||||
allowed: The new value.
|
||||
|
||||
Returns:
|
||||
bool: Whether a row was updated.
|
||||
"""
|
||||
t = sources_table
|
||||
result = self._conn.execute(
|
||||
t.update()
|
||||
.where(t.c.id == source_id)
|
||||
.where(t.c.user_id == user_id)
|
||||
.values(wiki_outside_edits=bool(allowed), updated_at=func.now())
|
||||
)
|
||||
return result.rowcount > 0
|
||||
|
||||
def get_by_legacy_id(
|
||||
self, legacy_mongo_id: str, user_id: Optional[str] = None,
|
||||
) -> Optional[dict]:
|
||||
|
||||
@@ -23,8 +23,11 @@ from docsgpt.storage.db.base_repository import looks_like_uuid, row_to_dict
|
||||
|
||||
|
||||
_JSONB_COLUMNS = {"config", "config_requirements", "actions"}
|
||||
_SCALAR_COLUMNS = {"name", "custom_name", "display_name", "description", "status"}
|
||||
_ALLOWED_COLUMNS = _SCALAR_COLUMNS | _JSONB_COLUMNS
|
||||
_SCALAR_COLUMNS = {"name", "custom_name", "display_name", "description", "status", "credential_mode"}
|
||||
# Set by server code only (tool creation, connection setup); route handlers
|
||||
# must not pass client input here.
|
||||
_UUID_COLUMNS = {"connection_id"}
|
||||
_ALLOWED_COLUMNS = _SCALAR_COLUMNS | _JSONB_COLUMNS | _UUID_COLUMNS
|
||||
|
||||
|
||||
def _encode_jsonb(value: Any) -> Any:
|
||||
@@ -60,6 +63,8 @@ class UserToolsRepository:
|
||||
status: bool = True,
|
||||
extra: Optional[dict] = None,
|
||||
legacy_mongo_id: Optional[str] = None,
|
||||
connection_id: Optional[str] = None,
|
||||
credential_mode: str = "owner",
|
||||
) -> dict:
|
||||
"""Insert a new tool row. ``extra`` is merged into the config JSONB."""
|
||||
cfg = config or {}
|
||||
@@ -70,14 +75,16 @@ class UserToolsRepository:
|
||||
"""
|
||||
INSERT INTO user_tools (
|
||||
user_id, name, custom_name, display_name, description,
|
||||
config, config_requirements, actions, status, legacy_mongo_id
|
||||
config, config_requirements, actions, status, legacy_mongo_id,
|
||||
connection_id, credential_mode
|
||||
)
|
||||
VALUES (
|
||||
:user_id, :name, :custom_name, :display_name, :description,
|
||||
CAST(:config AS jsonb),
|
||||
CAST(:config_requirements AS jsonb),
|
||||
CAST(:actions AS jsonb),
|
||||
:status, :legacy_mongo_id
|
||||
:status, :legacy_mongo_id,
|
||||
CAST(:connection_id AS uuid), :credential_mode
|
||||
)
|
||||
RETURNING *
|
||||
"""
|
||||
@@ -93,6 +100,8 @@ class UserToolsRepository:
|
||||
"actions": _encode_jsonb(actions or []),
|
||||
"status": status,
|
||||
"legacy_mongo_id": legacy_mongo_id,
|
||||
"connection_id": str(connection_id) if connection_id else None,
|
||||
"credential_mode": credential_mode,
|
||||
},
|
||||
)
|
||||
return row_to_dict(result.fetchone())
|
||||
@@ -213,6 +222,9 @@ class UserToolsRepository:
|
||||
if col in _JSONB_COLUMNS:
|
||||
set_clauses.append(f"{col} = CAST(:{col} AS jsonb)")
|
||||
params[col] = _encode_jsonb(val)
|
||||
elif col in _UUID_COLUMNS:
|
||||
set_clauses.append(f"{col} = CAST(:{col} AS uuid)")
|
||||
params[col] = str(val) if val else None
|
||||
else:
|
||||
set_clauses.append(f"{col} = :{col}")
|
||||
params[col] = val
|
||||
|
||||
@@ -20,6 +20,7 @@ from the repository itself rather than assumed.
|
||||
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import threading
|
||||
from dataclasses import replace
|
||||
from typing import Any, List, Optional
|
||||
@@ -39,6 +40,16 @@ _register_lock = threading.Lock()
|
||||
_FALLBACK_ONNX_FILE = "onnx/model.onnx"
|
||||
_FALLBACK_POOLING = "mean"
|
||||
|
||||
# The files FastEmbed fetches with every model besides its graph (see
|
||||
# ``ModelManagement.download_files_from_huggingface``).
|
||||
_FASTEMBED_SUPPORT_FILES = (
|
||||
"config.json",
|
||||
"tokenizer.json",
|
||||
"tokenizer_config.json",
|
||||
"special_tokens_map.json",
|
||||
"preprocessor_config.json",
|
||||
)
|
||||
|
||||
# Sentence-transformers records how a model turns token vectors into one
|
||||
# vector, and whether it normalises the result, as files in the repository.
|
||||
# Reading them is the difference between running a model and running something
|
||||
@@ -244,6 +255,51 @@ def _spec_for(model_name: str) -> EmbeddingModel:
|
||||
)
|
||||
|
||||
|
||||
def _complete_model_cache(repo: str, cache_dir: Optional[str]) -> None:
|
||||
"""Download a model's graph when its cached snapshot lacks it.
|
||||
|
||||
FastEmbed first loads a model from the cache and only downloads when that
|
||||
fails, and a snapshot of the repository counts as cached whatever it
|
||||
holds. The chunker caches only ``tokenizer.json`` there (see
|
||||
``docsgpt/parser/tokenization.py``), so FastEmbed then finds no graph and
|
||||
fails to load. Fetching the files it needs first makes either order work.
|
||||
Offline, nothing is fetched and FastEmbed reports what is missing; a
|
||||
failed download is logged and left to FastEmbed's own sources too.
|
||||
|
||||
Args:
|
||||
repo: The model's FastEmbed name.
|
||||
cache_dir: ``EMBEDDINGS_CACHE_DIR``, or None for FastEmbed's default.
|
||||
"""
|
||||
from fastembed import TextEmbedding
|
||||
from fastembed.common.utils import define_cache_dir
|
||||
from huggingface_hub import hf_hub_download, snapshot_download
|
||||
|
||||
lowered = repo.lower()
|
||||
description = next(
|
||||
(d for d in TextEmbedding._list_supported_models() if str(getattr(d, "model", "")).lower() == lowered),
|
||||
None,
|
||||
)
|
||||
source = getattr(getattr(description, "sources", None), "hf", None)
|
||||
model_file = getattr(description, "model_file", None)
|
||||
if not source or not isinstance(model_file, str):
|
||||
return # Not fetched from Hugging Face: FastEmbed handles it.
|
||||
needed = [model_file, *(getattr(description, "additional_files", None) or [])]
|
||||
cache = str(define_cache_dir(cache_dir))
|
||||
try:
|
||||
for filename in needed:
|
||||
hf_hub_download(source, filename, cache_dir=cache, local_files_only=True)
|
||||
return
|
||||
except Exception: # noqa: BLE001 -- any miss means the snapshot is incomplete
|
||||
pass
|
||||
if os.environ.get("HF_HUB_OFFLINE", "").strip().upper() in {"1", "TRUE", "YES", "ON"}:
|
||||
return
|
||||
logger.warning("The cached %s has no %s; downloading the model files.", source, ", ".join(needed))
|
||||
try:
|
||||
snapshot_download(repo_id=source, allow_patterns=[*_FASTEMBED_SUPPORT_FILES, *needed], cache_dir=cache)
|
||||
except Exception as exc: # noqa: BLE001 -- best effort: FastEmbed still tries its own sources
|
||||
logger.warning("Could not complete the cached %s (%s); leaving the download to FastEmbed.", source, exc)
|
||||
|
||||
|
||||
def _pad_to_longest_in_batch(model: Any) -> None:
|
||||
"""Undo a fixed padding width baked into a model's ``tokenizer.json``."""
|
||||
# FastEmbed enables padding only when the tokenizer declares none, so a
|
||||
@@ -295,6 +351,7 @@ class EmbeddingsWrapper:
|
||||
cache_dir = settings.EMBEDDINGS_CACHE_DIR
|
||||
if cache_dir:
|
||||
init_kwargs["cache_dir"] = cache_dir
|
||||
_complete_model_cache(self.spec.repo, cache_dir or None)
|
||||
self.model = TextEmbedding(**init_kwargs)
|
||||
except Exception as exc:
|
||||
raise RuntimeError(
|
||||
|
||||
+200
-10
@@ -32,6 +32,7 @@ from docsgpt.parser.file.image_parser import (
|
||||
VISION_CONVERTIBLE_MIME_TYPES,
|
||||
convert_image_to_png,
|
||||
)
|
||||
from docsgpt.parser.remote.github_loader import GitHubTokenRejected
|
||||
from docsgpt.parser.remote.remote_creator import (
|
||||
RemoteCreator,
|
||||
normalize_remote_data,
|
||||
@@ -1237,6 +1238,7 @@ def remote_worker(
|
||||
config=None,
|
||||
idempotency_key=None,
|
||||
source_id=None,
|
||||
connection_id=None,
|
||||
):
|
||||
safe_user = safe_filename(user)
|
||||
full_path = os.path.join(directory, safe_user, uuid.uuid4().hex)
|
||||
@@ -1283,7 +1285,26 @@ def remote_worker(
|
||||
self.update_state(state="PROGRESS", meta={"current": 1})
|
||||
logging.info("Initializing remote loader with type: %s", loader)
|
||||
remote_loader = RemoteCreator.create_loader(loader)
|
||||
raw_docs = remote_loader.load_data(source_data)
|
||||
loader_input = source_data
|
||||
if connection_id:
|
||||
loader_input = _with_connection_credentials(source_data, connection_id)
|
||||
if loader_input is None:
|
||||
from docsgpt.connectors.service import ConnectionUnavailable
|
||||
|
||||
raise ConnectionUnavailable("Reconnect to continue", connection_id=str(connection_id))
|
||||
try:
|
||||
raw_docs = remote_loader.load_data(loader_input)
|
||||
except GitHubTokenRejected as exc:
|
||||
# A revoked token pauses the connection's sources until the
|
||||
# owner reconnects, instead of failing on every schedule.
|
||||
from docsgpt.connectors import service as connection_service
|
||||
|
||||
if not connection_id:
|
||||
raise
|
||||
connection_service.mark_reconnect_needed(str(connection_id), str(exc))
|
||||
raise connection_service.ConnectionUnavailable(
|
||||
str(exc), connection_id=str(connection_id),
|
||||
) from exc
|
||||
|
||||
cfg = SourceConfig.parse(config)
|
||||
chunker = ChunkerCreator.create_chunker(
|
||||
@@ -1414,6 +1435,8 @@ def remote_worker(
|
||||
f"Failed to update last_sync for source {source_id_for_events}: {upd_err}"
|
||||
)
|
||||
upload_index(full_path, file_data)
|
||||
if connection_id:
|
||||
_link_source_to_connection(source_id_for_events, str(connection_id))
|
||||
publish_user_event(
|
||||
user,
|
||||
"source.ingest.completed",
|
||||
@@ -1467,6 +1490,7 @@ def sync(
|
||||
retriever,
|
||||
doc_id=None,
|
||||
directory="temp",
|
||||
connection_id=None,
|
||||
):
|
||||
try:
|
||||
remote_worker(
|
||||
@@ -1480,6 +1504,7 @@ def sync(
|
||||
sync_frequency,
|
||||
"sync",
|
||||
doc_id,
|
||||
connection_id=connection_id,
|
||||
)
|
||||
except Exception as e:
|
||||
logging.error(f"Error during sync: {e}", exc_info=True)
|
||||
@@ -1487,6 +1512,10 @@ def sync(
|
||||
return {"status": "success"}
|
||||
|
||||
|
||||
# Remote loaders that can only read with a connection (no public fallback).
|
||||
_CONNECTION_ONLY_LOADERS = frozenset({"linear"})
|
||||
|
||||
|
||||
def sync_worker(self, frequency):
|
||||
from sqlalchemy import text as sql_text
|
||||
|
||||
@@ -1494,7 +1523,7 @@ def sync_worker(self, frequency):
|
||||
with db_readonly() as conn:
|
||||
result = conn.execute(
|
||||
sql_text(
|
||||
"SELECT id, name, user_id, type, remote_data, retriever "
|
||||
"SELECT id, name, user_id, type, remote_data, retriever, connection_id, metadata "
|
||||
"FROM sources WHERE sync_frequency = :freq"
|
||||
),
|
||||
{"freq": frequency},
|
||||
@@ -1511,9 +1540,38 @@ def sync_worker(self, frequency):
|
||||
|
||||
sync_counts["total_sync_count"] += 1
|
||||
|
||||
# Connector sources have no RemoteCreator loader and need an OAuth
|
||||
# token to sync, which a scheduled task lacks — skip them.
|
||||
# Connector sources sync from their connection, whose token the
|
||||
# worker can refresh. Legacy ones with no connection still need the
|
||||
# browser, so they are skipped as before.
|
||||
if source_type and source_type.startswith("connector"):
|
||||
if doc.get("connection_id"):
|
||||
from docsgpt.api.user.tasks import sync_connector_source as sync_task
|
||||
|
||||
sync_task.delay(doc_id)
|
||||
sync_counts["sync_dispatched"] += 1
|
||||
else:
|
||||
sync_counts["sync_skipped"] += 1
|
||||
continue
|
||||
|
||||
metadata = doc.get("metadata")
|
||||
if isinstance(metadata, str):
|
||||
try:
|
||||
metadata = json.loads(metadata)
|
||||
except ValueError:
|
||||
metadata = {}
|
||||
if (
|
||||
doc.get("connection_id")
|
||||
and isinstance(metadata, dict)
|
||||
and metadata.get("sync_state") == "paused_reconnect"
|
||||
):
|
||||
# An S3 or GitHub source whose connection needs reconnecting:
|
||||
# it resumes when the owner reconnects, rather than failing (and
|
||||
# notifying) on every schedule until then.
|
||||
sync_counts["sync_skipped"] += 1
|
||||
continue
|
||||
if source_type in _CONNECTION_ONLY_LOADERS and not doc.get("connection_id"):
|
||||
# Linear is read only with a connection's sign-in; its
|
||||
# connection was removed and the content kept.
|
||||
sync_counts["sync_skipped"] += 1
|
||||
continue
|
||||
|
||||
@@ -1525,7 +1583,8 @@ def sync_worker(self, frequency):
|
||||
continue
|
||||
|
||||
resp = sync(
|
||||
self, source_data, name, user, source_type, frequency, retriever, doc_id
|
||||
self, source_data, name, user, source_type, frequency, retriever, doc_id,
|
||||
connection_id=str(doc["connection_id"]) if doc.get("connection_id") else None,
|
||||
)
|
||||
sync_counts[
|
||||
"sync_success" if resp["status"] == "success" else "sync_failure"
|
||||
@@ -1533,7 +1592,7 @@ def sync_worker(self, frequency):
|
||||
return {
|
||||
key: sync_counts[key]
|
||||
for key in [
|
||||
"total_sync_count", "sync_success", "sync_failure", "sync_skipped",
|
||||
"total_sync_count", "sync_success", "sync_failure", "sync_skipped", "sync_dispatched",
|
||||
]
|
||||
}
|
||||
|
||||
@@ -2161,6 +2220,132 @@ def _webhook_tool_allowlist(agent_config):
|
||||
return []
|
||||
|
||||
|
||||
def _with_connection_credentials(source_data, connection_id: str):
|
||||
"""Loader input with the connection's stored keys merged in, or None.
|
||||
|
||||
S3, Reddit and GitHub sources made from a connection keep their keys on
|
||||
the connection only, never in ``sources.remote_data``. A JSON string
|
||||
stays a JSON string and a dict a dict; any other string (a GitHub
|
||||
repository URL) becomes ``{"url": ...}`` next to the keys. An MCP
|
||||
sign-in (Linear) gets its ``connection_id`` instead: its tokens stay
|
||||
with the MCP client, which renews them. Returns None when the
|
||||
connection is gone, needs reconnecting, or its connector is turned off.
|
||||
"""
|
||||
from docsgpt.connectors import service
|
||||
from docsgpt.storage.db.repositories.connector_sessions import ConnectorSessionsRepository
|
||||
|
||||
with db_readonly() as conn:
|
||||
row = ConnectorSessionsRepository(conn).get(str(connection_id))
|
||||
enabled = row is not None and service.connector_enabled(conn, row)
|
||||
if row is None or not enabled:
|
||||
return None
|
||||
if (row.get("auth_kind") or "") == "mcp_oauth":
|
||||
if service.normalize_status(row) != service.STATUS_CONNECTED:
|
||||
return None
|
||||
credentials = {"connection_id": str(row["id"])}
|
||||
else:
|
||||
try:
|
||||
credentials = service.access_credentials(row)
|
||||
except service.ConnectionUnavailable:
|
||||
return None
|
||||
if isinstance(source_data, str):
|
||||
try:
|
||||
parsed = json.loads(source_data)
|
||||
except ValueError:
|
||||
parsed = None
|
||||
if isinstance(parsed, dict):
|
||||
return json.dumps({**parsed, **credentials})
|
||||
return {"url": source_data, **credentials}
|
||||
return {**dict(source_data or {}), **credentials}
|
||||
|
||||
|
||||
def _link_source_to_connection(source_id: str, connection_id: str) -> None:
|
||||
"""Point a source at the connection it syncs from, and lift any reconnect pause."""
|
||||
from sqlalchemy import text as sql_text
|
||||
|
||||
try:
|
||||
with db_session() as conn:
|
||||
conn.execute(
|
||||
sql_text(
|
||||
"UPDATE sources SET connection_id = CAST(:cid AS uuid), "
|
||||
"metadata = metadata - 'sync_state' WHERE id = CAST(:sid AS uuid)"
|
||||
),
|
||||
{"cid": str(connection_id), "sid": str(source_id)},
|
||||
)
|
||||
except Exception:
|
||||
logging.warning("Could not link source %s to connection %s", source_id, connection_id, exc_info=True)
|
||||
|
||||
|
||||
def sync_connector_source(self, source_id: str) -> Dict[str, Any]:
|
||||
"""Re-download and re-index a connector source from its connection.
|
||||
|
||||
Runs as the connection owner with no browser: the connection service
|
||||
refreshes the token under a row lock. When the grant was revoked the
|
||||
service flags the connection and pauses its sources, and this returns
|
||||
``paused`` rather than failing again on every schedule.
|
||||
|
||||
Args:
|
||||
self: The bound Celery task.
|
||||
source_id: The source to sync.
|
||||
|
||||
Returns:
|
||||
``{"status": "success" | "paused" | "disabled" | "skipped"}`` plus
|
||||
the ingest result.
|
||||
"""
|
||||
from docsgpt.connectors.service import ConnectionUnavailable, connector_enabled, normalize_status
|
||||
from docsgpt.storage.db.repositories.connector_sessions import ConnectorSessionsRepository
|
||||
|
||||
with db_readonly() as conn:
|
||||
from sqlalchemy import text as sql_text
|
||||
|
||||
row = conn.execute(
|
||||
sql_text(
|
||||
"SELECT id, name, user_id, remote_data, retriever, sync_frequency, config, connection_id "
|
||||
"FROM sources WHERE id = CAST(:id AS uuid)"
|
||||
),
|
||||
{"id": str(source_id)},
|
||||
).fetchone()
|
||||
source = dict(row._mapping) if row else None
|
||||
connection = (
|
||||
ConnectorSessionsRepository(conn).get(str(source["connection_id"]))
|
||||
if source and source.get("connection_id")
|
||||
else None
|
||||
)
|
||||
enabled = connection is not None and connector_enabled(conn, connection)
|
||||
if not source or not connection:
|
||||
return {"status": "skipped"}
|
||||
if not enabled:
|
||||
return {"status": "disabled"}
|
||||
if normalize_status(connection) != "connected":
|
||||
return {"status": "paused"}
|
||||
remote_data = source.get("remote_data") or {}
|
||||
if isinstance(remote_data, str):
|
||||
try:
|
||||
remote_data = json.loads(remote_data)
|
||||
except json.JSONDecodeError:
|
||||
remote_data = {}
|
||||
provider = remote_data.get("provider") or connection.get("provider")
|
||||
try:
|
||||
result = ingest_connector(
|
||||
self,
|
||||
source.get("name"),
|
||||
source.get("user_id"),
|
||||
provider,
|
||||
connection_id=str(connection["id"]),
|
||||
file_ids=remote_data.get("file_ids") or [],
|
||||
folder_ids=remote_data.get("folder_ids") or [],
|
||||
recursive=remote_data.get("recursive", True),
|
||||
retriever=source.get("retriever") or "classic",
|
||||
operation_mode="sync",
|
||||
doc_id=str(source["id"]),
|
||||
sync_frequency=source.get("sync_frequency") or "never",
|
||||
config=source.get("config") or None,
|
||||
)
|
||||
except ConnectionUnavailable:
|
||||
return {"status": "paused"}
|
||||
return {"status": "success", "result": result}
|
||||
|
||||
|
||||
def ingest_connector(
|
||||
self,
|
||||
job_name: str,
|
||||
@@ -2177,6 +2362,7 @@ def ingest_connector(
|
||||
config=None,
|
||||
idempotency_key=None,
|
||||
source_id=None,
|
||||
connection_id=None,
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Ingestion for internal knowledge bases (GoogleDrive, etc.).
|
||||
@@ -2185,7 +2371,8 @@ def ingest_connector(
|
||||
job_name: Name of the ingestion job
|
||||
user: User identifier
|
||||
source_type: Type of remote source ("google_drive", "dropbox", etc.)
|
||||
session_token: Authentication token for the service
|
||||
session_token: Legacy browser session token naming the connection
|
||||
connection_id: The connection whose account the files are read with
|
||||
file_ids: List of file IDs to download
|
||||
folder_ids: List of folder IDs to download
|
||||
recursive: Whether to recursively download folders
|
||||
@@ -2247,8 +2434,8 @@ def ingest_connector(
|
||||
meta={"current": 10, "status": "Initializing connector"},
|
||||
)
|
||||
|
||||
if not session_token:
|
||||
raise ValueError(f"{source_type} connector requires session_token")
|
||||
if not session_token and not connection_id:
|
||||
raise ValueError(f"{source_type} connector requires a connection")
|
||||
|
||||
if not ConnectorCreator.is_supported(source_type):
|
||||
raise ValueError(
|
||||
@@ -2256,8 +2443,9 @@ def ingest_connector(
|
||||
)
|
||||
|
||||
remote_loader = ConnectorCreator.create_connector(
|
||||
source_type, session_token
|
||||
source_type, session_token, connection_id=connection_id
|
||||
)
|
||||
connection_id = remote_loader.connection_id
|
||||
|
||||
# Create a clean config for storage
|
||||
api_source_config = {
|
||||
@@ -2411,6 +2599,8 @@ def ingest_connector(
|
||||
)
|
||||
|
||||
upload_index(vector_store_path, file_data)
|
||||
if connection_id:
|
||||
_link_source_to_connection(source_id_for_events, connection_id)
|
||||
|
||||
# Ensure we mark the task as complete
|
||||
self.update_state(
|
||||
|
||||
+13
-8
@@ -647,10 +647,12 @@ description use a Switch in a SettingRow instead.
|
||||
### Switch, TimePicker and Calendar (`ui/switch.tsx`, `ui/time-picker.tsx`, `ui/calendar.tsx`)
|
||||
|
||||
`Switch` (Radix) is an on/off setting, placed inside a `SettingRow` that names
|
||||
it. A tile's own on/off (the tool tile's "In my chats") is the one exception:
|
||||
a `Label text-muted-foreground text-xs font-normal` + `Switch` pair at the
|
||||
tile's bottom-right. For a shared tool it is the caller's own preference,
|
||||
never a switch that turns the tool off for everyone. The track is `primary` when on and `bg-input` when off, with a white
|
||||
it. A tool's own "In my chats" on/off is the one exception. On a Tools page
|
||||
tile it is a bare `Switch` at the tile's bottom-right, named only by its
|
||||
`aria-label` (no visible label); in a connection's drawer, a tool row ends in
|
||||
a `Label text-muted-foreground text-xs font-normal` + `Switch` pair. Either way
|
||||
it is the caller's own preference, never a switch that turns the tool off for
|
||||
everyone. The track is `primary` when on and `bg-input` when off, with a white
|
||||
thumb in both themes (see Elevation).
|
||||
|
||||
`TimePicker` picks a time of day as two `SelectTrigger`s, hours and minutes,
|
||||
@@ -806,9 +808,10 @@ overlay. The app has one
|
||||
`ToastViewport`, mounted in `App.tsx`; it is the live region
|
||||
(`role="status"`, `aria-live="polite"`) and the fixed stack, so `Toast`
|
||||
cards carry no role and no toast renders its own rail or positioning. Top
|
||||
to bottom it holds `TeamNotificationToast`, `ToolApprovalToast`,
|
||||
`UploadToast` and `ActionToast`, and it moves to the bottom-left while any
|
||||
agent preview drawer is open (workflow or classic). A new toast component returns only its
|
||||
to bottom it holds `TeamNotificationToast`, `ConnectionHealthToast` (a
|
||||
connection that needs reconnecting, with a Reconnect action),
|
||||
`ToolApprovalToast`, `UploadToast` and `ActionToast`, and it moves to the
|
||||
bottom-left while any agent preview drawer is open (workflow or classic). A new toast component returns only its
|
||||
`Toast` cards and is added to that viewport. A page that reports the result
|
||||
of an action (the admin Users actions) dispatches
|
||||
`showActionToast({ variant: 'success' | 'destructive', message })` from
|
||||
@@ -888,7 +891,8 @@ also refuses Logs, Schedules and Pin on a draft), never on `ownership`,
|
||||
(`components/ViewOnlyNotice`: an `Alert role="note"` with `Lock` first and
|
||||
`common.viewOnlyNotice`) as its first child. Where only part of an editable
|
||||
form is locked (a tool's credentials when the owner turned off "Editors can
|
||||
change credentials"), the same notice passes its own `message`. There is no Save; Cancel becomes a lone Close.
|
||||
change credentials", or when the tool runs on the owner's connection, whose
|
||||
secret only the owner changes), the same notice passes its own `message`. There is no Save; Cancel becomes a lone Close.
|
||||
Fields are `disabled` (a group: `<fieldset disabled className="min-w-0">`);
|
||||
long text the user reads or copies (a prompt, a chunk) is `readOnly`, so it
|
||||
keeps full contrast and scrolls. Titles and menu items swap Edit and
|
||||
@@ -1670,6 +1674,7 @@ list stays reviewable.
|
||||
| `settings/PersonalAccessTokens.tsx` | `shadcn/no-restyle` | Token scope chips are identifiers, so the `neutral` Badge is set in mono (`font-mono`). One disable. |
|
||||
| `agents/workflow/WorkflowPreview.tsx` | `shadcn/no-restyle` | A step's state changes (keys and values the app serialised) are `neutral` Badges set in mono (`font-mono`), like token scopes. One disable. |
|
||||
| `settings/traces/TraceChips.tsx` | `shadcn/no-restyle` | Trace stat chips (durations, counts) use tabular figures so they don't jitter between rows (`tabular-nums` on Badge). One disable. |
|
||||
| `connectors/ConnectorSetupNotice.tsx` | `shadcn/no-restyle` | The setup-guide link inside the needs-setup warning Alert keeps the Alert's status colour (`text-current` on `link inline`). One disable. |
|
||||
| `modals/MCPServerModal.tsx` | `shadcn/no-restyle` | The authorization link inside the test-result Alert keeps the Alert's status colour (`text-current` on `link inline`). One disable. |
|
||||
| `components/MermaidRenderer.tsx` | `shadcn/no-restyle` | The zoom readout between − and + is a `link inline` Button on the `bg-black/70` overlay; it keeps the overlay's white 12px regular text (`text-xs font-normal text-current`). One disable, beside the two zoom-button ones above. |
|
||||
| `conversation/SharedConversation.tsx` | `shadcn/no-restyle` | The "DocsGPT" link sits in the `/share/:id` page's regular-weight byline (`font-normal`). One disable. |
|
||||
|
||||
@@ -43,6 +43,7 @@ import Notification from './components/Notification';
|
||||
import ToolApprovalToast from './notifications/ToolApprovalToast';
|
||||
import TeamNotificationToast from './notifications/TeamNotificationToast';
|
||||
import ActionToast from './notifications/ActionToast';
|
||||
import ConnectionHealthToast from './notifications/ConnectionHealthToast';
|
||||
|
||||
function AuthWrapper({ children }: { children: React.ReactNode }) {
|
||||
const { t } = useTranslation();
|
||||
@@ -128,6 +129,7 @@ function MainLayout() {
|
||||
onMouseDown={(e) => e.stopPropagation()}
|
||||
>
|
||||
<TeamNotificationToast />
|
||||
<ConnectionHealthToast />
|
||||
<ToolApprovalToast />
|
||||
<UploadToast />
|
||||
<ActionToast />
|
||||
|
||||
@@ -0,0 +1,121 @@
|
||||
import { configureStore } from '@reduxjs/toolkit';
|
||||
import { act } from 'react';
|
||||
import { createRoot, type Root } from 'react-dom/client';
|
||||
import { Provider } from 'react-redux';
|
||||
import { MemoryRouter, Route, Routes } from 'react-router-dom';
|
||||
|
||||
vi.mock('react-i18next', () => ({
|
||||
useTranslation: () => ({
|
||||
t: (key: string, opts?: Record<string, unknown>) => {
|
||||
if (key === 'demo')
|
||||
return [1, 2, 3, 4].map((n) => ({
|
||||
header: `demo${n}`,
|
||||
query: `q${n}`,
|
||||
}));
|
||||
return opts?.names ? `${key}:${opts.names}` : key;
|
||||
},
|
||||
}),
|
||||
}));
|
||||
|
||||
vi.mock('./api/services/modelService', () => ({
|
||||
default: {
|
||||
getModels: vi.fn(async () => ({ ok: true, json: async () => ({}) })),
|
||||
transformModels: () => [],
|
||||
},
|
||||
}));
|
||||
vi.mock('./hooks', () => ({ useDarkTheme: () => [false] }));
|
||||
|
||||
import connectorsReducer from './connectors/connectorsSlice';
|
||||
import Hero from './Hero';
|
||||
|
||||
Object.assign(globalThis, { IS_REACT_ACT_ENVIRONMENT: true });
|
||||
|
||||
const connector = (key: string, name: string, publisher = 'built_in') => ({
|
||||
key,
|
||||
name,
|
||||
publisher,
|
||||
available: true,
|
||||
});
|
||||
|
||||
describe('Hero connect card', () => {
|
||||
let container: HTMLDivElement;
|
||||
let root: Root;
|
||||
|
||||
beforeEach(() => {
|
||||
container = document.createElement('div');
|
||||
document.body.appendChild(container);
|
||||
root = createRoot(container);
|
||||
});
|
||||
|
||||
afterEach(async () => {
|
||||
await act(async () => root.unmount());
|
||||
container.remove();
|
||||
});
|
||||
|
||||
const render = async (connections: unknown[]) => {
|
||||
const store = configureStore({
|
||||
reducer: {
|
||||
connectors: connectorsReducer,
|
||||
preference: (
|
||||
state = {
|
||||
token: null,
|
||||
selectedModel: null,
|
||||
availableModels: [],
|
||||
modelsLoading: false,
|
||||
},
|
||||
) => state,
|
||||
},
|
||||
preloadedState: {
|
||||
connectors: {
|
||||
enabled: true,
|
||||
loading: false,
|
||||
loaded: true,
|
||||
failed: false,
|
||||
catalog: [
|
||||
connector('google_drive', 'Google Drive'),
|
||||
connector('mcp:notion', 'Notion', 'preset'),
|
||||
connector('custom_mcp', 'MCP server', 'custom'),
|
||||
],
|
||||
connections,
|
||||
},
|
||||
},
|
||||
} as Parameters<typeof configureStore>[0]);
|
||||
await act(async () => {
|
||||
root.render(
|
||||
<Provider store={store}>
|
||||
<MemoryRouter initialEntries={['/']}>
|
||||
<Routes>
|
||||
<Route path="/" element={<Hero handleQuestion={vi.fn()} />} />
|
||||
<Route
|
||||
path="/settings/connectors"
|
||||
element={<div>CONNECTORS</div>}
|
||||
/>
|
||||
</Routes>
|
||||
</MemoryRouter>
|
||||
</Provider>,
|
||||
);
|
||||
});
|
||||
};
|
||||
|
||||
const card = () =>
|
||||
container.querySelector<HTMLButtonElement>(
|
||||
'[data-testid="hero-connect-card"]',
|
||||
);
|
||||
|
||||
it('offers connecting the services this install has', async () => {
|
||||
await render([]);
|
||||
expect(card()?.textContent).toContain(
|
||||
'connectHero.body:Google Drive, Notion',
|
||||
);
|
||||
// It replaces the last demo card.
|
||||
expect(container.textContent).not.toContain('demo4');
|
||||
await act(async () => card()!.click());
|
||||
expect(container.textContent).toContain('CONNECTORS');
|
||||
});
|
||||
|
||||
it('keeps the demo cards once something is connected', async () => {
|
||||
await render([{ id: 'c1', status: 'connected' }]);
|
||||
expect(card()).toBeNull();
|
||||
expect(container.textContent).toContain('demo4');
|
||||
});
|
||||
});
|
||||
+55
-1
@@ -1,5 +1,6 @@
|
||||
import { useEffect, useRef } from 'react';
|
||||
import { useTranslation } from 'react-i18next';
|
||||
import { useNavigate } from 'react-router-dom';
|
||||
import { useDispatch, useSelector } from 'react-redux';
|
||||
|
||||
import modelService from './api/services/modelService';
|
||||
@@ -14,6 +15,13 @@ import {
|
||||
SelectTrigger,
|
||||
SelectValue,
|
||||
} from './components/ui/select';
|
||||
import {
|
||||
selectConnections,
|
||||
selectConnectorCatalog,
|
||||
selectConnectorsEnabled,
|
||||
selectConnectorsLoaded,
|
||||
} from './connectors/connectorsSlice';
|
||||
import { connectorName } from './connectors/i18n';
|
||||
import { useDarkTheme } from './hooks';
|
||||
import {
|
||||
selectAvailableModels,
|
||||
@@ -145,10 +153,31 @@ export default function Hero({
|
||||
}) {
|
||||
const { t } = useTranslation();
|
||||
const [isDarkTheme] = useDarkTheme();
|
||||
const navigate = useNavigate();
|
||||
const demos = t('demo', { returnObjects: true }) as Array<{
|
||||
header: string;
|
||||
query: string;
|
||||
}>;
|
||||
// Nothing connected yet: the last card offers connecting a service, named
|
||||
// after what this install can actually connect.
|
||||
const connectorsEnabled = useSelector(selectConnectorsEnabled);
|
||||
const connectorsLoaded = useSelector(selectConnectorsLoaded);
|
||||
const connections = useSelector(selectConnections);
|
||||
const catalog = useSelector(selectConnectorCatalog);
|
||||
const connectable = catalog.filter(
|
||||
(connector) => connector.available && connector.publisher !== 'custom',
|
||||
);
|
||||
const offerConnect =
|
||||
connectorsEnabled &&
|
||||
connectorsLoaded &&
|
||||
connections.length === 0 &&
|
||||
connectable.length > 0;
|
||||
const connectNames = connectable
|
||||
.slice(0, 2)
|
||||
.map((connector) => connectorName(t, connector))
|
||||
.join(', ');
|
||||
const cards = (demos ?? []).filter((demo) => demo.header && demo.query);
|
||||
const shown = offerConnect ? cards.slice(0, 3) : cards;
|
||||
|
||||
return (
|
||||
<div className="text-foreground flex h-full w-full flex-col items-center justify-between">
|
||||
@@ -170,7 +199,7 @@ export default function Hero({
|
||||
{/* Demo Buttons Section */}
|
||||
<div className="mb-3 w-full max-w-full md:mb-3">
|
||||
<div className="grid grid-cols-1 gap-3 text-xs md:grid-cols-1 md:gap-4 lg:grid-cols-2">
|
||||
{demos?.map(
|
||||
{shown.map(
|
||||
(demo: { header: string; query: string }, key: number) =>
|
||||
demo.header &&
|
||||
demo.query && (
|
||||
@@ -199,6 +228,31 @@ export default function Hero({
|
||||
</Button>
|
||||
),
|
||||
)}
|
||||
{offerConnect && (
|
||||
<Button
|
||||
type="button"
|
||||
variant="outline"
|
||||
size="lg"
|
||||
shape="pill"
|
||||
onClick={() => navigate('/settings/connectors')}
|
||||
data-testid="hero-connect-card"
|
||||
className={cn(
|
||||
/* eslint-disable-next-line shadcn/no-restyle --
|
||||
Same two-line pill as the demo cards above. */
|
||||
'hidden h-auto w-full flex-col items-start gap-0 py-3.5 text-left text-xs font-normal whitespace-normal md:flex',
|
||||
)}
|
||||
>
|
||||
<p className="text-foreground mb-2 font-semibold">
|
||||
{t('connectHero.title')}
|
||||
</p>
|
||||
<span className="text-muted-foreground line-clamp-2">
|
||||
{t('connectHero.body', {
|
||||
names: connectNames,
|
||||
interpolation: { escapeValue: false },
|
||||
})}
|
||||
</span>
|
||||
</Button>
|
||||
)}
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
@@ -56,6 +56,7 @@ import {
|
||||
getSectionForPath,
|
||||
type Section,
|
||||
} from './navigation/sections';
|
||||
import ConnectionHealthDot from './connectors/ConnectionHealthDot';
|
||||
import { useSidebarLevel } from './navigation/SidebarLevelProvider';
|
||||
import { useSectionContext } from './navigation/useSectionContext';
|
||||
import { useLastAppPath } from './navigation/useLastAppPath';
|
||||
@@ -798,6 +799,7 @@ export default function Navigation({ navOpen, setNavOpen }: NavigationProps) {
|
||||
aria-hidden
|
||||
/>
|
||||
<p className="text-foreground text-sm">{t('settings.label')}</p>
|
||||
<ConnectionHealthDot />
|
||||
</Link>
|
||||
</Button>
|
||||
</div>
|
||||
|
||||
@@ -0,0 +1,358 @@
|
||||
import { configureStore } from '@reduxjs/toolkit';
|
||||
import { act } from 'react';
|
||||
import { createRoot, type Root } from 'react-dom/client';
|
||||
import { Provider } from 'react-redux';
|
||||
|
||||
vi.mock('react-i18next', () => ({
|
||||
useTranslation: () => ({ t: (key: string) => key }),
|
||||
}));
|
||||
|
||||
const getAdmin = vi.fn();
|
||||
const updateAdmin = vi.fn();
|
||||
vi.mock('../api/services/connectorsService', () => ({
|
||||
default: {
|
||||
getAdmin: (...args: unknown[]) => getAdmin(...args),
|
||||
updateAdmin: (...args: unknown[]) => updateAdmin(...args),
|
||||
},
|
||||
}));
|
||||
|
||||
import actionToastReducer, {
|
||||
selectActionToast,
|
||||
} from '../notifications/actionToastSlice';
|
||||
import { prefSlice } from '../preferences/preferenceSlice';
|
||||
import Connectors from './Connectors';
|
||||
|
||||
Object.assign(globalThis, { IS_REACT_ACT_ENVIRONMENT: true });
|
||||
|
||||
const connector = (overrides: Record<string, unknown> = {}) => ({
|
||||
key: 'google_drive',
|
||||
name: 'Google Drive',
|
||||
icon: 'google-drive',
|
||||
publisher: 'built_in',
|
||||
auth_kind: 'oauth',
|
||||
capabilities: ['sync'],
|
||||
enabled: true,
|
||||
credential_mode: 'choose',
|
||||
configured: false,
|
||||
required_settings: [
|
||||
{ name: 'GOOGLE_CLIENT_ID', set: true },
|
||||
{ name: 'GOOGLE_CLIENT_SECRET', set: false },
|
||||
],
|
||||
connection_count: 3,
|
||||
docs_url: null,
|
||||
mcp_url: null,
|
||||
...overrides,
|
||||
});
|
||||
|
||||
const MCP_ROW = connector({
|
||||
key: 'custom_mcp',
|
||||
name: 'MCP server',
|
||||
icon: 'mcp',
|
||||
publisher: 'custom',
|
||||
auth_kind: 'mcp',
|
||||
capabilities: ['read', 'write'],
|
||||
configured: true,
|
||||
required_settings: [],
|
||||
connection_count: 0,
|
||||
});
|
||||
const NOTION = {
|
||||
key: 'mcp_notion',
|
||||
name: 'Notion',
|
||||
icon: 'notion',
|
||||
publisher: 'preset',
|
||||
auth_kind: 'mcp_oauth',
|
||||
capabilities: ['read', 'write'],
|
||||
configured: true,
|
||||
required_settings: [],
|
||||
connection_count: 0,
|
||||
};
|
||||
|
||||
const payload = (overrides: Record<string, unknown> = {}) => ({
|
||||
success: true,
|
||||
connectors: [
|
||||
connector(),
|
||||
connector({
|
||||
key: 'mcp_notion',
|
||||
name: 'Notion',
|
||||
icon: 'notion',
|
||||
publisher: 'preset',
|
||||
auth_kind: 'mcp_oauth',
|
||||
capabilities: ['read', 'write'],
|
||||
configured: true,
|
||||
required_settings: [],
|
||||
connection_count: 0,
|
||||
}),
|
||||
MCP_ROW,
|
||||
],
|
||||
allow_custom_mcp: true,
|
||||
default_encryption_key: false,
|
||||
oauth_redirect_uri: 'https://docs.example/api/connectors/callback',
|
||||
mcp_redirect_uri: 'https://docs.example/api/mcp_server/callback',
|
||||
...overrides,
|
||||
});
|
||||
|
||||
describe('Admin Connectors', () => {
|
||||
let container: HTMLDivElement;
|
||||
let root: Root;
|
||||
|
||||
beforeEach(() => {
|
||||
getAdmin.mockReset();
|
||||
updateAdmin.mockReset();
|
||||
container = document.createElement('div');
|
||||
document.body.appendChild(container);
|
||||
root = createRoot(container);
|
||||
});
|
||||
|
||||
afterEach(() => {
|
||||
act(() => root.unmount());
|
||||
container.remove();
|
||||
});
|
||||
|
||||
let store: ReturnType<typeof makeStore>;
|
||||
const makeStore = () =>
|
||||
configureStore({
|
||||
reducer: {
|
||||
preference: prefSlice.reducer,
|
||||
actionToast: actionToastReducer,
|
||||
},
|
||||
});
|
||||
|
||||
const render = async () => {
|
||||
store = makeStore();
|
||||
await act(async () => {
|
||||
root.render(
|
||||
<Provider store={store}>
|
||||
<Connectors />
|
||||
</Provider>,
|
||||
);
|
||||
});
|
||||
};
|
||||
|
||||
it('lists each connector with its setup state and redirect URIs', async () => {
|
||||
getAdmin.mockResolvedValue(payload());
|
||||
await render();
|
||||
const rows = Array.from(container.querySelectorAll('tbody tr'));
|
||||
expect(rows).toHaveLength(3);
|
||||
expect(rows[0].textContent).toContain('Google Drive');
|
||||
expect(rows[0].textContent).toContain('Needs setup');
|
||||
expect(rows[0].textContent).toContain('Setup guide');
|
||||
expect(rows[1].textContent).toContain('Preset');
|
||||
expect(rows[1].textContent).toContain('Ready');
|
||||
expect(rows[1].textContent).not.toContain('Setup guide');
|
||||
expect(container.textContent).toContain(
|
||||
'https://docs.example/api/connectors/callback',
|
||||
);
|
||||
});
|
||||
|
||||
it('shows GitHub as ready, with its optional GitHub App settings in the guide', async () => {
|
||||
const GITHUB = connector({
|
||||
key: 'github',
|
||||
name: 'GitHub',
|
||||
icon: 'github',
|
||||
auth_kind: 'api_key',
|
||||
capabilities: ['sync', 'read'],
|
||||
configured: true,
|
||||
required_settings: [],
|
||||
oauth_settings: [
|
||||
{ name: 'GITHUB_CLIENT_ID', set: true },
|
||||
{ name: 'GITHUB_CLIENT_SECRET', set: false },
|
||||
{ name: 'GITHUB_APP_SLUG', set: false },
|
||||
],
|
||||
oauth_configured: false,
|
||||
connection_count: 0,
|
||||
});
|
||||
getAdmin.mockResolvedValue(payload({ connectors: [GITHUB] }));
|
||||
await render();
|
||||
const row = container.querySelector('tbody tr')!;
|
||||
expect(row.textContent).toContain('Ready');
|
||||
expect(row.textContent).toContain('Tokens only');
|
||||
const guide = Array.from(row.querySelectorAll('button')).find(
|
||||
(b) => b.textContent === 'Setup guide',
|
||||
)!;
|
||||
await act(async () => guide.click());
|
||||
const text = document.body.textContent ?? '';
|
||||
expect(text).toContain('Sign in with GitHub');
|
||||
expect(text).toContain('GITHUB_APP_SLUG');
|
||||
expect(text).toContain(
|
||||
'Request user authorization (OAuth) during installation',
|
||||
);
|
||||
expect(text).toContain('https://docs.example/api/connectors/callback');
|
||||
});
|
||||
|
||||
it('warns when credentials use the default encryption key', async () => {
|
||||
getAdmin.mockResolvedValue(payload({ default_encryption_key: true }));
|
||||
await render();
|
||||
expect(container.textContent).toContain(
|
||||
'Set ENCRYPTION_SECRET_KEY before connecting services.',
|
||||
);
|
||||
expect(container.textContent).toContain('docsgpt connectors reencrypt');
|
||||
});
|
||||
|
||||
it('saves a connector toggle as a policy', async () => {
|
||||
getAdmin.mockResolvedValue(payload());
|
||||
updateAdmin.mockResolvedValue(
|
||||
payload({ connectors: [connector({ ...NOTION, enabled: false })] }),
|
||||
);
|
||||
await render();
|
||||
const toggle = () =>
|
||||
container.querySelector<HTMLButtonElement>(
|
||||
'table [aria-label="Notion enabled"]',
|
||||
)!;
|
||||
await act(async () => toggle().click());
|
||||
expect(updateAdmin).toHaveBeenCalledWith(
|
||||
{ policies: { mcp_notion: { enabled: false } } },
|
||||
null,
|
||||
);
|
||||
expect(toggle().getAttribute('aria-checked')).toBe('false');
|
||||
});
|
||||
|
||||
it('keeps a connector that needs setup off, and says why', async () => {
|
||||
getAdmin.mockResolvedValue(payload());
|
||||
await render();
|
||||
const drive = container.querySelector<HTMLButtonElement>(
|
||||
'table [aria-label="Google Drive enabled"]',
|
||||
)!;
|
||||
expect(drive.disabled).toBe(true);
|
||||
});
|
||||
|
||||
it('shows no sharing policy for a sync-only connector', async () => {
|
||||
getAdmin.mockResolvedValue(payload());
|
||||
await render();
|
||||
const [drive, notion] = Array.from(container.querySelectorAll('tbody tr'));
|
||||
expect(drive.textContent).toContain('No tools');
|
||||
expect(
|
||||
notion.querySelector('[aria-label="Notion sharing policy"]'),
|
||||
).not.toBeNull();
|
||||
});
|
||||
|
||||
it('turns custom MCP servers off from their own row', async () => {
|
||||
getAdmin.mockResolvedValue(payload());
|
||||
updateAdmin.mockResolvedValue(payload({ allow_custom_mcp: false }));
|
||||
await render();
|
||||
const mcp = () =>
|
||||
container.querySelector<HTMLButtonElement>(
|
||||
'table [aria-label="MCP server enabled"]',
|
||||
)!;
|
||||
await act(async () => mcp().click());
|
||||
expect(updateAdmin).toHaveBeenCalledWith(
|
||||
{ allow_custom_mcp: false, policies: { custom_mcp: { enabled: false } } },
|
||||
null,
|
||||
);
|
||||
expect(mcp().getAttribute('aria-checked')).toBe('false');
|
||||
expect(container.querySelector('#allow-custom-mcp')).toBeNull();
|
||||
});
|
||||
|
||||
it('lists connectors on phones and opens their controls in a sheet', async () => {
|
||||
getAdmin.mockResolvedValue(payload());
|
||||
await render();
|
||||
const rows = Array.from(
|
||||
container.querySelectorAll<HTMLButtonElement>(
|
||||
'[data-slot="list-row"] button',
|
||||
),
|
||||
);
|
||||
expect(rows).toHaveLength(3);
|
||||
expect(rows[0].textContent).toContain('Off · 3 connections · No tools');
|
||||
await act(async () => rows[1].click());
|
||||
expect(
|
||||
document.body.querySelector('[data-slot="sheet-content"]'),
|
||||
).not.toBeNull();
|
||||
expect(document.body.textContent).toContain('Shared tools use');
|
||||
});
|
||||
|
||||
it('reports a failed save in a toast and keeps the page', async () => {
|
||||
getAdmin.mockResolvedValue(payload());
|
||||
updateAdmin.mockResolvedValue({ success: false });
|
||||
await render();
|
||||
await act(async () =>
|
||||
container
|
||||
.querySelector<HTMLButtonElement>(
|
||||
'table [aria-label="Notion enabled"]',
|
||||
)!
|
||||
.click(),
|
||||
);
|
||||
expect(updateAdmin).toHaveBeenCalledWith(
|
||||
{ policies: { mcp_notion: { enabled: false } } },
|
||||
null,
|
||||
);
|
||||
expect(selectActionToast(store.getState())).toMatchObject({
|
||||
variant: 'destructive',
|
||||
message: 'Could not save the change.',
|
||||
});
|
||||
expect(container.querySelectorAll('tbody tr')).toHaveLength(3);
|
||||
});
|
||||
|
||||
it('saves one change at a time so a late response cannot win', async () => {
|
||||
getAdmin.mockResolvedValue(payload());
|
||||
const pending: ((value: unknown) => void)[] = [];
|
||||
updateAdmin.mockImplementation(
|
||||
() => new Promise((resolve) => pending.push(resolve)),
|
||||
);
|
||||
await render();
|
||||
const notion = () =>
|
||||
container.querySelector<HTMLButtonElement>(
|
||||
'table [aria-label="Notion enabled"]',
|
||||
)!;
|
||||
const mcp = () =>
|
||||
container.querySelector<HTMLButtonElement>(
|
||||
'table [aria-label="MCP server enabled"]',
|
||||
)!;
|
||||
await act(async () => notion().click());
|
||||
await act(async () => mcp().click());
|
||||
// The second save waits for the first.
|
||||
expect(updateAdmin).toHaveBeenCalledTimes(1);
|
||||
const off = connector({ ...NOTION, enabled: false });
|
||||
await act(async () => pending[0](payload({ connectors: [off, MCP_ROW] })));
|
||||
expect(updateAdmin).toHaveBeenCalledTimes(2);
|
||||
await act(async () =>
|
||||
pending[1](
|
||||
payload({ connectors: [off, MCP_ROW], allow_custom_mcp: false }),
|
||||
),
|
||||
);
|
||||
expect(notion().getAttribute('aria-checked')).toBe('false');
|
||||
expect(mcp().getAttribute('aria-checked')).toBe('false');
|
||||
});
|
||||
|
||||
describe('write access', () => {
|
||||
const GITHUB = connector({
|
||||
key: 'github',
|
||||
name: 'GitHub',
|
||||
icon: 'github',
|
||||
auth_kind: 'api_key',
|
||||
capabilities: ['sync', 'read'],
|
||||
configured: true,
|
||||
required_settings: [],
|
||||
allow_writes: true,
|
||||
});
|
||||
const writes = () =>
|
||||
container.querySelector<HTMLButtonElement>('#allow-writes-github');
|
||||
|
||||
it('lets an admin forbid changes through GitHub', async () => {
|
||||
getAdmin.mockResolvedValue(payload({ connectors: [GITHUB, MCP_ROW] }));
|
||||
updateAdmin.mockResolvedValue(
|
||||
payload({
|
||||
connectors: [{ ...GITHUB, allow_writes: false }, MCP_ROW],
|
||||
}),
|
||||
);
|
||||
await render();
|
||||
expect(writes()!.getAttribute('aria-checked')).toBe('true');
|
||||
await act(async () => writes()!.click());
|
||||
expect(updateAdmin).toHaveBeenCalledWith(
|
||||
{ policies: { github: { allow_writes: false } } },
|
||||
null,
|
||||
);
|
||||
expect(writes()!.getAttribute('aria-checked')).toBe('false');
|
||||
});
|
||||
|
||||
it('has no write access section without such a connector', async () => {
|
||||
getAdmin.mockResolvedValue(payload());
|
||||
await render();
|
||||
expect(container.textContent).not.toContain('Write access');
|
||||
});
|
||||
});
|
||||
|
||||
it('offers a retry when loading fails', async () => {
|
||||
getAdmin.mockRejectedValue(new Error('offline'));
|
||||
await render();
|
||||
expect(container.textContent).toContain('Failed to load connectors.');
|
||||
});
|
||||
});
|
||||
@@ -0,0 +1,609 @@
|
||||
import { ExternalLink, Info, TriangleAlert } from 'lucide-react';
|
||||
import { useCallback, useEffect, useRef, useState } from 'react';
|
||||
import { useDispatch, useSelector } from 'react-redux';
|
||||
|
||||
import connectorsService from '../api/services/connectorsService';
|
||||
import CopyButton from '../components/CopyButton';
|
||||
import PageToolbar from '../components/PageToolbar';
|
||||
import { Alert, AlertDescription, AlertTitle } from '../components/ui/alert';
|
||||
import { Badge } from '../components/ui/badge';
|
||||
import { Button } from '../components/ui/button';
|
||||
import { Card } from '../components/ui/card';
|
||||
import { FormField } from '../components/ui/form-field';
|
||||
import { ListRow, ListRows } from '../components/ui/list-row';
|
||||
import {
|
||||
Sheet,
|
||||
SheetContent,
|
||||
SheetDescription,
|
||||
SheetTitle,
|
||||
} from '../components/ui/sheet';
|
||||
import {
|
||||
DescriptionItem,
|
||||
DescriptionList,
|
||||
} from '../components/ui/description-list';
|
||||
import { LoadingState } from '../components/ui/loading-state';
|
||||
import { Modal } from '../components/ui/modal';
|
||||
import { SectionHeader } from '../components/ui/section-header';
|
||||
import {
|
||||
Select,
|
||||
SelectContent,
|
||||
SelectItem,
|
||||
SelectTrigger,
|
||||
SelectValue,
|
||||
} from '../components/ui/select';
|
||||
import { SettingRow, SettingRows } from '../components/ui/setting-row';
|
||||
import { Switch } from '../components/ui/switch';
|
||||
import {
|
||||
Table,
|
||||
TableBody,
|
||||
TableCell,
|
||||
TableContainer,
|
||||
TableHead,
|
||||
TableHeader,
|
||||
TableRow,
|
||||
} from '../components/ui/table';
|
||||
import {
|
||||
Tooltip,
|
||||
TooltipContent,
|
||||
TooltipTrigger,
|
||||
} from '../components/ui/tooltip';
|
||||
import ConnectorIcon from '../connectors/ConnectorIcon';
|
||||
import { showActionToast } from '../notifications/actionToastSlice';
|
||||
import { selectToken } from '../preferences/preferenceSlice';
|
||||
import { LoadError, fmtNumber } from './AdminUI';
|
||||
|
||||
type Policy = 'choose' | 'owner' | 'member';
|
||||
|
||||
type AdminConnector = {
|
||||
key: string;
|
||||
name: string;
|
||||
icon: string;
|
||||
publisher: 'built_in' | 'preset' | 'custom';
|
||||
auth_kind: string;
|
||||
capabilities: string[];
|
||||
enabled: boolean;
|
||||
credential_mode: Policy;
|
||||
configured: boolean;
|
||||
required_settings: { name: string; set: boolean }[];
|
||||
/** Optional settings that add a second sign-in (GitHub's App). */
|
||||
oauth_settings?: { name: string; set: boolean }[];
|
||||
oauth_configured?: boolean;
|
||||
connection_count: number;
|
||||
docs_url: string | null;
|
||||
mcp_url: string | null;
|
||||
/**
|
||||
* Whether members may let agents make changes through it (GitHub's write
|
||||
* tools); null where the connector offers no such choice.
|
||||
*/
|
||||
allow_writes?: boolean | null;
|
||||
};
|
||||
|
||||
type AdminConnectorsData = {
|
||||
success: boolean;
|
||||
connectors: AdminConnector[];
|
||||
allow_custom_mcp: boolean;
|
||||
default_encryption_key: boolean;
|
||||
oauth_redirect_uri: string;
|
||||
mcp_redirect_uri: string;
|
||||
};
|
||||
|
||||
const POLICY_LABELS: Record<Policy, string> = {
|
||||
choose: 'The sharer decides per share',
|
||||
member: "Always each person's own account",
|
||||
owner: "Always the sharer's account",
|
||||
};
|
||||
|
||||
const hasTools = (connector: AdminConnector) =>
|
||||
connector.capabilities.some((capability) => capability !== 'sync');
|
||||
|
||||
const hasSetupGuide = (connector: AdminConnector) =>
|
||||
connector.required_settings.length > 0 ||
|
||||
(connector.oauth_settings?.length ?? 0) > 0;
|
||||
|
||||
/** Works with pasted tokens, but its optional OAuth sign-in is not set up yet. */
|
||||
const tokensOnly = (connector: AdminConnector) =>
|
||||
(connector.oauth_settings?.length ?? 0) > 0 && !connector.oauth_configured;
|
||||
|
||||
function SettingsList({
|
||||
settings,
|
||||
}: {
|
||||
settings: { name: string; set: boolean }[];
|
||||
}) {
|
||||
return (
|
||||
<DescriptionList layout="justified" size="xs">
|
||||
{settings.map((setting) => (
|
||||
<DescriptionItem
|
||||
key={setting.name}
|
||||
label={<code className="font-mono">{setting.name}</code>}
|
||||
>
|
||||
<Badge variant={setting.set ? 'success' : 'warning'}>
|
||||
{setting.set ? 'Set' : 'Missing'}
|
||||
</Badge>
|
||||
</DescriptionItem>
|
||||
))}
|
||||
</DescriptionList>
|
||||
);
|
||||
}
|
||||
|
||||
function CodeRow({ value }: { value: string }) {
|
||||
return (
|
||||
<Card variant="filled" padding="sm" className="flex-row items-start gap-2">
|
||||
<pre className="min-w-0 flex-1 font-mono text-xs wrap-anywhere whitespace-pre-wrap">
|
||||
{value}
|
||||
</pre>
|
||||
<CopyButton textToCopy={value} />
|
||||
</Card>
|
||||
);
|
||||
}
|
||||
|
||||
function SetupGuide({
|
||||
connector,
|
||||
redirectUri,
|
||||
onClose,
|
||||
}: {
|
||||
connector: AdminConnector;
|
||||
redirectUri: string;
|
||||
onClose: () => void;
|
||||
}) {
|
||||
const oauthSettings = connector.oauth_settings ?? [];
|
||||
const optionalOAuth =
|
||||
connector.required_settings.length === 0 && oauthSettings.length > 0;
|
||||
return (
|
||||
<Modal
|
||||
open
|
||||
onOpenChange={(open) => !open && onClose()}
|
||||
title={`Set up ${connector.name}`}
|
||||
description={
|
||||
optionalOAuth
|
||||
? `Members can already connect ${connector.name} with their own access tokens. To also offer Sign in with ${connector.name}, register a GitHub App, then set these server settings and restart the API and the worker.`
|
||||
: 'Register DocsGPT as an OAuth app with the provider, then set these server settings and restart the API and the worker.'
|
||||
}
|
||||
footer={
|
||||
<Button size="lg" shape="pill" onClick={onClose}>
|
||||
Done
|
||||
</Button>
|
||||
}
|
||||
>
|
||||
<div className="flex flex-col gap-6">
|
||||
<section className="flex flex-col gap-2">
|
||||
<SectionHeader
|
||||
as="h3"
|
||||
size="xs"
|
||||
title={
|
||||
optionalOAuth
|
||||
? 'Callback URL to register'
|
||||
: 'Redirect URI to register'
|
||||
}
|
||||
/>
|
||||
<CodeRow value={redirectUri} />
|
||||
</section>
|
||||
{connector.required_settings.length > 0 && (
|
||||
<section className="flex flex-col gap-2">
|
||||
<SectionHeader as="h3" size="xs" title="Server settings" />
|
||||
<SettingsList settings={connector.required_settings} />
|
||||
</section>
|
||||
)}
|
||||
{oauthSettings.length > 0 && (
|
||||
<section className="flex flex-col gap-2">
|
||||
<SectionHeader
|
||||
as="h3"
|
||||
size="xs"
|
||||
title={`Sign in with ${connector.name} (optional)`}
|
||||
/>
|
||||
<SettingsList settings={oauthSettings} />
|
||||
</section>
|
||||
)}
|
||||
{connector.key === 'github' && (
|
||||
<Alert variant="info" role="note">
|
||||
<Info />
|
||||
<AlertDescription>
|
||||
In the GitHub App, give repository permissions Contents and
|
||||
Metadata read-only access, and turn on Request user authorization
|
||||
(OAuth) during installation so choosing repositories returns to
|
||||
DocsGPT. For agents to make changes, also give Issues and Pull
|
||||
requests read and write access (Contents read and write only if
|
||||
agents should edit files). GITHUB_APP_SLUG is the name in the
|
||||
app's public link (github.com/apps/<slug>).
|
||||
</AlertDescription>
|
||||
</Alert>
|
||||
)}
|
||||
{connector.key === 'google_drive' && (
|
||||
<Alert variant="info" role="note">
|
||||
<Info />
|
||||
<AlertDescription>
|
||||
Publish the Google OAuth app (or use an internal Workspace app).
|
||||
Apps left in Testing get refresh tokens that expire after seven
|
||||
days, which stops background sync.
|
||||
</AlertDescription>
|
||||
</Alert>
|
||||
)}
|
||||
{connector.docs_url && (
|
||||
<Button variant="link" size="inline" asChild className="w-fit">
|
||||
<a
|
||||
href={connector.docs_url}
|
||||
target="_blank"
|
||||
rel="noopener noreferrer"
|
||||
>
|
||||
Setup guide
|
||||
<ExternalLink />
|
||||
</a>
|
||||
</Button>
|
||||
)}
|
||||
</div>
|
||||
</Modal>
|
||||
);
|
||||
}
|
||||
|
||||
/**
|
||||
* Admin > Connectors: which connectors members may use, whose account a
|
||||
* shared tool runs with, and what each OAuth connector still needs from the
|
||||
* server. English only, like the rest of the admin pages.
|
||||
*/
|
||||
export default function Connectors() {
|
||||
const dispatch = useDispatch();
|
||||
const token = useSelector(selectToken);
|
||||
const [data, setData] = useState<AdminConnectorsData | null>(null);
|
||||
const [loading, setLoading] = useState(true);
|
||||
const [guide, setGuide] = useState<AdminConnector | null>(null);
|
||||
const [detailKey, setDetailKey] = useState<string | null>(null);
|
||||
|
||||
const load = useCallback(async () => {
|
||||
setLoading(true);
|
||||
try {
|
||||
setData(await connectorsService.getAdmin(token));
|
||||
} catch {
|
||||
setData({ success: false } as AdminConnectorsData);
|
||||
} finally {
|
||||
setLoading(false);
|
||||
}
|
||||
}, [token]);
|
||||
|
||||
useEffect(() => {
|
||||
load();
|
||||
}, [load]);
|
||||
|
||||
// Saves run one after another: each response is a full snapshot, so an
|
||||
// older one arriving last would otherwise put back a stale policy.
|
||||
const saveQueue = useRef<Promise<void>>(Promise.resolve());
|
||||
const save = (body: {
|
||||
policies?: Record<
|
||||
string,
|
||||
{ enabled?: boolean; credential_mode?: Policy; allow_writes?: boolean }
|
||||
>;
|
||||
allow_custom_mcp?: boolean;
|
||||
}) => {
|
||||
const failed = () =>
|
||||
dispatch(
|
||||
showActionToast({
|
||||
variant: 'destructive',
|
||||
message: 'Could not save the change.',
|
||||
}),
|
||||
);
|
||||
saveQueue.current = saveQueue.current.then(async () => {
|
||||
try {
|
||||
const next = await connectorsService.updateAdmin(body, token);
|
||||
if (next?.success) setData(next);
|
||||
else failed();
|
||||
} catch {
|
||||
failed();
|
||||
}
|
||||
});
|
||||
return saveQueue.current;
|
||||
};
|
||||
|
||||
if (data === null && loading) return <LoadingState fill="block" />;
|
||||
if (!data?.success)
|
||||
return <LoadError message="Failed to load connectors." onRetry={load} />;
|
||||
|
||||
// The sheet (phones) always shows the connector's latest saved state.
|
||||
const detail = data.connectors.find((c) => c.key === detailKey) ?? null;
|
||||
const writable = data.connectors.filter(
|
||||
(connector) =>
|
||||
connector.allow_writes !== null && connector.allow_writes !== undefined,
|
||||
);
|
||||
|
||||
// The custom MCP row is the one switch for members' own MCP servers.
|
||||
const isCustomMcp = (connector: AdminConnector) =>
|
||||
connector.key === 'custom_mcp';
|
||||
const enabledOf = (connector: AdminConnector) =>
|
||||
connector.enabled && (!isCustomMcp(connector) || data.allow_custom_mcp);
|
||||
const setEnabled = (connector: AdminConnector, on: boolean) =>
|
||||
save(
|
||||
isCustomMcp(connector)
|
||||
? { allow_custom_mcp: on, policies: { custom_mcp: { enabled: on } } }
|
||||
: { policies: { [connector.key]: { enabled: on } } },
|
||||
);
|
||||
|
||||
const statusBadge = (connector: AdminConnector) => (
|
||||
<span className="inline-flex flex-wrap gap-1">
|
||||
<Badge variant={connector.configured ? 'success' : 'warning'}>
|
||||
{connector.configured ? 'Ready' : 'Needs setup'}
|
||||
</Badge>
|
||||
{connector.configured && tokensOnly(connector) && (
|
||||
<Badge variant="neutral">Tokens only</Badge>
|
||||
)}
|
||||
</span>
|
||||
);
|
||||
|
||||
const summary = (connector: AdminConnector) =>
|
||||
[
|
||||
// The badge beside it says whether setup is missing; this is what
|
||||
// members get.
|
||||
connector.configured && enabledOf(connector) ? 'On' : 'Off',
|
||||
`${fmtNumber(connector.connection_count)} ${
|
||||
connector.connection_count === 1 ? 'connection' : 'connections'
|
||||
}`,
|
||||
hasTools(connector)
|
||||
? POLICY_LABELS[connector.credential_mode]
|
||||
: 'No tools',
|
||||
].join(' · ');
|
||||
|
||||
const enabledSwitch = (connector: AdminConnector) => {
|
||||
const control = (
|
||||
<Switch
|
||||
checked={enabledOf(connector)}
|
||||
disabled={!connector.configured}
|
||||
aria-label={`${connector.name} enabled`}
|
||||
onCheckedChange={(checked) => setEnabled(connector, checked === true)}
|
||||
/>
|
||||
);
|
||||
// Off until its server settings exist: switching it on would do nothing.
|
||||
if (connector.configured) return control;
|
||||
return (
|
||||
<Tooltip>
|
||||
<TooltipTrigger asChild>
|
||||
<span tabIndex={0} className="inline-flex w-fit">
|
||||
{control}
|
||||
</span>
|
||||
</TooltipTrigger>
|
||||
<TooltipContent>Add its server settings first</TooltipContent>
|
||||
</Tooltip>
|
||||
);
|
||||
};
|
||||
|
||||
const policyControl = (connector: AdminConnector, fullWidth = false) =>
|
||||
hasTools(connector) ? (
|
||||
<Select
|
||||
value={connector.credential_mode}
|
||||
onValueChange={(value) =>
|
||||
save({
|
||||
policies: { [connector.key]: { credential_mode: value as Policy } },
|
||||
})
|
||||
}
|
||||
>
|
||||
<SelectTrigger
|
||||
size="sm"
|
||||
className={fullWidth ? 'w-full' : 'w-60'}
|
||||
aria-label={`${connector.name} sharing policy`}
|
||||
>
|
||||
<SelectValue />
|
||||
</SelectTrigger>
|
||||
<SelectContent>
|
||||
{(Object.keys(POLICY_LABELS) as Policy[]).map((policy) => (
|
||||
<SelectItem key={policy} value={policy}>
|
||||
{POLICY_LABELS[policy]}
|
||||
</SelectItem>
|
||||
))}
|
||||
</SelectContent>
|
||||
</Select>
|
||||
) : (
|
||||
<span className="text-muted-foreground text-sm">No tools</span>
|
||||
);
|
||||
|
||||
return (
|
||||
<div className="flex flex-col gap-8">
|
||||
<PageToolbar intro="Choose which connectors members can use and whose account a shared tool runs with. A connector that still needs server settings starts turned off and is hidden from members; it turns on once its settings are in place, unless you switch it off." />
|
||||
|
||||
{data.default_encryption_key && (
|
||||
<Alert variant="destructive">
|
||||
<TriangleAlert />
|
||||
<AlertTitle>
|
||||
Set ENCRYPTION_SECRET_KEY before connecting services.
|
||||
</AlertTitle>
|
||||
<AlertDescription>
|
||||
Stored credentials are encrypted with ENCRYPTION_SECRET_KEY, which
|
||||
still has its public default value. Set your own, keep the old one
|
||||
in ENCRYPTION_SECRET_KEY_PREVIOUS and run{' '}
|
||||
<code className="font-mono text-xs">
|
||||
docsgpt connectors reencrypt
|
||||
</code>
|
||||
.
|
||||
</AlertDescription>
|
||||
</Alert>
|
||||
)}
|
||||
|
||||
<section className="flex flex-col gap-3">
|
||||
<SectionHeader
|
||||
title="Redirect URIs"
|
||||
description="Register these with each provider's OAuth app."
|
||||
/>
|
||||
<div className="grid grid-cols-1 gap-x-4 gap-y-3 lg:grid-cols-2">
|
||||
<div className="flex flex-col gap-1.5">
|
||||
<span className="text-muted-foreground text-xs">
|
||||
Google Drive, SharePoint, Confluence and the GitHub App
|
||||
</span>
|
||||
<CodeRow value={data.oauth_redirect_uri} />
|
||||
</div>
|
||||
<div className="flex flex-col gap-1.5">
|
||||
<span className="text-muted-foreground text-xs">MCP servers</span>
|
||||
<CodeRow value={data.mcp_redirect_uri} />
|
||||
</div>
|
||||
</div>
|
||||
</section>
|
||||
|
||||
<section className="flex flex-col gap-3">
|
||||
<SectionHeader
|
||||
title="Connectors"
|
||||
description="The custom MCP server row decides whether members can add their own MCP servers; presets are switched one by one."
|
||||
/>
|
||||
{/* Phones: a list; each row opens the connector's controls. */}
|
||||
<Card padding="none" className="overflow-hidden md:hidden">
|
||||
<ListRows>
|
||||
{data.connectors.map((connector) => (
|
||||
<ListRow
|
||||
key={connector.key}
|
||||
interactive
|
||||
asChild
|
||||
leading={
|
||||
<ConnectorIcon
|
||||
icon={connector.icon}
|
||||
className="size-5 shrink-0"
|
||||
/>
|
||||
}
|
||||
title={connector.name}
|
||||
description={summary(connector)}
|
||||
trailing={statusBadge(connector)}
|
||||
>
|
||||
<button
|
||||
type="button"
|
||||
onClick={() => setDetailKey(connector.key)}
|
||||
/>
|
||||
</ListRow>
|
||||
))}
|
||||
</ListRows>
|
||||
</Card>
|
||||
<TableContainer className="hidden md:block">
|
||||
<Table>
|
||||
<TableHead>
|
||||
<TableRow>
|
||||
<TableHeader>Connector</TableHeader>
|
||||
<TableHeader>Status</TableHeader>
|
||||
<TableHeader align="right">Connections</TableHeader>
|
||||
<TableHeader>Enabled</TableHeader>
|
||||
<TableHeader>Shared tools use</TableHeader>
|
||||
<TableHeader />
|
||||
</TableRow>
|
||||
</TableHead>
|
||||
<TableBody>
|
||||
{data.connectors.map((connector) => (
|
||||
<TableRow key={connector.key}>
|
||||
<TableCell>
|
||||
<span className="flex items-center gap-2">
|
||||
<ConnectorIcon
|
||||
icon={connector.icon}
|
||||
className="size-5 shrink-0"
|
||||
/>
|
||||
<span className="truncate">{connector.name}</span>
|
||||
{connector.publisher !== 'built_in' && (
|
||||
<Badge variant="neutral">
|
||||
{connector.publisher === 'preset'
|
||||
? 'Preset'
|
||||
: 'Custom'}
|
||||
</Badge>
|
||||
)}
|
||||
</span>
|
||||
</TableCell>
|
||||
<TableCell>{statusBadge(connector)}</TableCell>
|
||||
<TableCell align="right" className="tabular-nums">
|
||||
{fmtNumber(connector.connection_count)}
|
||||
</TableCell>
|
||||
<TableCell>{enabledSwitch(connector)}</TableCell>
|
||||
<TableCell>{policyControl(connector)}</TableCell>
|
||||
<TableCell align="right">
|
||||
{hasSetupGuide(connector) && (
|
||||
<Button
|
||||
type="button"
|
||||
variant="outline"
|
||||
size="xs"
|
||||
onClick={() => setGuide(connector)}
|
||||
>
|
||||
Setup guide
|
||||
</Button>
|
||||
)}
|
||||
</TableCell>
|
||||
</TableRow>
|
||||
))}
|
||||
</TableBody>
|
||||
</Table>
|
||||
</TableContainer>
|
||||
</section>
|
||||
|
||||
{writable.length > 0 && (
|
||||
<section className="flex flex-col gap-3">
|
||||
<SectionHeader
|
||||
title="Write access"
|
||||
description="Members can let agents make changes through these connectors, one connection at a time. Each change asks first unless the member allows it."
|
||||
/>
|
||||
<SettingRows>
|
||||
{writable.map((connector) => (
|
||||
<SettingRow
|
||||
key={connector.key}
|
||||
label={`Let agents make changes through ${connector.name}`}
|
||||
description={
|
||||
connector.key === 'github'
|
||||
? 'Create issues, comments and pull requests. Off keeps every GitHub tool read-only, including ones already set up for changes.'
|
||||
: `Off keeps every ${connector.name} tool read-only.`
|
||||
}
|
||||
htmlFor={`allow-writes-${connector.key}`}
|
||||
alignStart
|
||||
>
|
||||
<Switch
|
||||
id={`allow-writes-${connector.key}`}
|
||||
checked={connector.allow_writes === true}
|
||||
onCheckedChange={(checked) =>
|
||||
save({
|
||||
policies: {
|
||||
[connector.key]: { allow_writes: checked === true },
|
||||
},
|
||||
})
|
||||
}
|
||||
/>
|
||||
</SettingRow>
|
||||
))}
|
||||
</SettingRows>
|
||||
</section>
|
||||
)}
|
||||
|
||||
{detail && (
|
||||
<Sheet open onOpenChange={(open) => !open && setDetailKey(null)}>
|
||||
<SheetContent side="right" size="detail" closeLabel="Close">
|
||||
<div className="flex flex-col gap-6 p-6">
|
||||
<div className="flex items-center gap-3 pr-12">
|
||||
<ConnectorIcon icon={detail.icon} className="size-7" />
|
||||
<SheetTitle className="truncate">{detail.name}</SheetTitle>
|
||||
</div>
|
||||
<SheetDescription>{summary(detail)}</SheetDescription>
|
||||
<SettingRows>
|
||||
<SettingRow
|
||||
label="Enabled"
|
||||
description={
|
||||
detail.configured
|
||||
? 'Members can connect and use it.'
|
||||
: 'Add its server settings first (Setup guide).'
|
||||
}
|
||||
>
|
||||
{enabledSwitch(detail)}
|
||||
</SettingRow>
|
||||
</SettingRows>
|
||||
{hasTools(detail) && (
|
||||
<FormField label="Shared tools use">
|
||||
{policyControl(detail, true)}
|
||||
</FormField>
|
||||
)}
|
||||
{hasSetupGuide(detail) && (
|
||||
<Button
|
||||
type="button"
|
||||
variant="outline"
|
||||
size="sm"
|
||||
shape="pill"
|
||||
className="w-fit"
|
||||
onClick={() => setGuide(detail)}
|
||||
>
|
||||
Setup guide
|
||||
</Button>
|
||||
)}
|
||||
</div>
|
||||
</SheetContent>
|
||||
</Sheet>
|
||||
)}
|
||||
|
||||
{guide && (
|
||||
<SetupGuide
|
||||
connector={guide}
|
||||
redirectUri={data.oauth_redirect_uri}
|
||||
onClose={() => setGuide(null)}
|
||||
/>
|
||||
)}
|
||||
</div>
|
||||
);
|
||||
}
|
||||
@@ -8,6 +8,7 @@ import Admins from './Admins';
|
||||
import Activity from './Activity';
|
||||
import Overview from './Overview';
|
||||
import Quotas from './Quotas';
|
||||
import Connectors from './Connectors';
|
||||
import Usage from './Usage';
|
||||
import Users from './Users';
|
||||
|
||||
@@ -39,6 +40,7 @@ export default function Admin() {
|
||||
<Route path="roles" element={<Admins />} />
|
||||
<Route path="usage" element={<Usage />} />
|
||||
<Route path="quotas" element={<Quotas />} />
|
||||
<Route path="connectors" element={<Connectors />} />
|
||||
<Route path="audit" element={<Activity />} />
|
||||
<Route path="*" element={<Navigate to="/admin" replace />} />
|
||||
</Routes>
|
||||
|
||||
@@ -0,0 +1,444 @@
|
||||
import { configureStore } from '@reduxjs/toolkit';
|
||||
import { act } from 'react';
|
||||
import { createRoot, type Root } from 'react-dom/client';
|
||||
import { Provider } from 'react-redux';
|
||||
|
||||
vi.mock('react-i18next', () => ({
|
||||
useTranslation: () => ({
|
||||
t: (key: string, opts?: Record<string, unknown>) => {
|
||||
if (!opts) return key;
|
||||
const params = Object.entries(opts)
|
||||
.filter(([k]) => k !== 'interpolation' && k !== 'count')
|
||||
.map(([k, v]) => `${k}=${v}`)
|
||||
.join(',');
|
||||
return params ? `${key}(${params})` : key;
|
||||
},
|
||||
i18n: { language: 'en' },
|
||||
}),
|
||||
}));
|
||||
|
||||
const getUserTools = vi.fn();
|
||||
const getAgent = vi.fn();
|
||||
const getWorkflow = vi.fn();
|
||||
const updateAgent = vi.fn();
|
||||
vi.mock('../api/services/userService', () => ({
|
||||
default: {
|
||||
getUserTools: (...args: unknown[]) => getUserTools(...args),
|
||||
getAgent: (...args: unknown[]) => getAgent(...args),
|
||||
getWorkflow: (...args: unknown[]) => getWorkflow(...args),
|
||||
updateAgent: (...args: unknown[]) => updateAgent(...args),
|
||||
},
|
||||
}));
|
||||
|
||||
import actionToastReducer, {
|
||||
selectActionToast,
|
||||
} from '../notifications/actionToastSlice';
|
||||
import { prefSlice } from '../preferences/preferenceSlice';
|
||||
import ApiWriteAllowlist from './ApiWriteAllowlist';
|
||||
import type { Mock } from 'vitest';
|
||||
|
||||
import type { Agent, AgentConfig, ResourceState } from './types';
|
||||
|
||||
Object.assign(globalThis, { IS_REACT_ACT_ENVIRONMENT: true });
|
||||
|
||||
const TOOLS = {
|
||||
tools: [
|
||||
{
|
||||
id: 'tg',
|
||||
displayName: 'Telegram',
|
||||
connection_id: 'c1',
|
||||
owner_credential_writes: ['telegram_send_message'],
|
||||
actions: [
|
||||
{
|
||||
name: 'telegram_send_message',
|
||||
description: 'Sends a message.',
|
||||
access: 'write',
|
||||
active: true,
|
||||
},
|
||||
{ name: 'telegram_read', access: 'read', active: true },
|
||||
],
|
||||
},
|
||||
{
|
||||
id: 'memory',
|
||||
displayName: 'Memory',
|
||||
connection_id: null,
|
||||
owner_credential_writes: [],
|
||||
actions: [{ name: 'memory_write', access: 'write', active: true }],
|
||||
},
|
||||
{
|
||||
id: 'crm',
|
||||
displayName: 'CRM API',
|
||||
connection_id: null,
|
||||
owner_credential_writes: ['create_lead'],
|
||||
actions: [],
|
||||
},
|
||||
],
|
||||
};
|
||||
|
||||
// What the agent read says each tool may write on stored credentials: the
|
||||
// owner's own tools, one an editor sponsored that the owner can't see, and
|
||||
// one that stopped.
|
||||
const state = (over: Partial<ResourceState>): ResourceState => ({
|
||||
key: `tool:${over.id}`,
|
||||
type: 'tool',
|
||||
id: 'x',
|
||||
name: 'Tool',
|
||||
state: 'active',
|
||||
reason: null,
|
||||
owner_credential_writes: [],
|
||||
...over,
|
||||
});
|
||||
|
||||
const STATES: ResourceState[] = [
|
||||
state({
|
||||
id: 'tg',
|
||||
name: 'Telegram',
|
||||
owner_credential_writes: ['telegram_send_message', 'telegram_pin_message'],
|
||||
}),
|
||||
state({ id: 'memory', name: 'Memory' }),
|
||||
state({
|
||||
id: 'crm',
|
||||
name: 'CRM API',
|
||||
owner_credential_writes: ['create_lead'],
|
||||
}),
|
||||
state({
|
||||
id: 'bob-jira',
|
||||
name: 'Bob Jira',
|
||||
runs_as: { user_id: 'bob', label: 'bob@example.com' },
|
||||
owner_credential_writes: ['create_issue'],
|
||||
}),
|
||||
state({
|
||||
id: 'gone',
|
||||
name: 'Gone',
|
||||
state: 'stopped',
|
||||
reason: 'deleted',
|
||||
owner_credential_writes: ['delete_all'],
|
||||
}),
|
||||
];
|
||||
|
||||
const readAgent = (tools: string[], extra: Partial<Agent> = {}) => ({
|
||||
ok: true,
|
||||
json: async () => ({
|
||||
id: 'agent-1',
|
||||
agent_type: 'classic',
|
||||
tools,
|
||||
resource_states: STATES.filter((s) => tools.includes(s.id)),
|
||||
...extra,
|
||||
}),
|
||||
});
|
||||
|
||||
const agent = (overrides: Partial<Agent> = {}): Agent =>
|
||||
({
|
||||
id: 'agent-1',
|
||||
tools: ['tg', 'memory'],
|
||||
config: { guardrails: { controls: [] } },
|
||||
...overrides,
|
||||
}) as unknown as Agent;
|
||||
|
||||
describe('ApiWriteAllowlist', () => {
|
||||
let container: HTMLDivElement;
|
||||
let root: Root;
|
||||
let store: ReturnType<typeof makeStore>;
|
||||
const makeStore = () =>
|
||||
configureStore({
|
||||
reducer: {
|
||||
preference: prefSlice.reducer,
|
||||
actionToast: actionToastReducer,
|
||||
},
|
||||
});
|
||||
|
||||
beforeEach(() => {
|
||||
getUserTools.mockResolvedValue({ json: async () => TOOLS });
|
||||
getAgent
|
||||
.mockReset()
|
||||
.mockImplementation(async () => readAgent(currentTools));
|
||||
getWorkflow.mockReset();
|
||||
updateAgent.mockReset();
|
||||
container = document.createElement('div');
|
||||
document.body.appendChild(container);
|
||||
root = createRoot(container);
|
||||
});
|
||||
|
||||
afterEach(async () => {
|
||||
await act(async () => root.unmount());
|
||||
container.remove();
|
||||
});
|
||||
|
||||
let currentTools: string[] = [];
|
||||
const render = async (
|
||||
a: Agent,
|
||||
{
|
||||
onConfigChange = vi.fn<(config: AgentConfig) => void>(),
|
||||
getSavedConfig,
|
||||
defaultOpen,
|
||||
}: {
|
||||
onConfigChange?: Mock<(config: AgentConfig) => void>;
|
||||
getSavedConfig?: () => Agent['config'];
|
||||
defaultOpen?: boolean;
|
||||
} = {},
|
||||
) => {
|
||||
currentTools = a.tools ?? [];
|
||||
store = makeStore();
|
||||
await act(async () => {
|
||||
root.render(
|
||||
<Provider store={store}>
|
||||
<ApiWriteAllowlist
|
||||
agent={a}
|
||||
onConfigChange={onConfigChange}
|
||||
getSavedConfig={getSavedConfig}
|
||||
defaultOpen={defaultOpen}
|
||||
/>
|
||||
</Provider>,
|
||||
);
|
||||
});
|
||||
return onConfigChange;
|
||||
};
|
||||
|
||||
const K = 'modals.agentDetails.apiWrites';
|
||||
const disclosure = () =>
|
||||
Array.from(container.querySelectorAll('button')).find((b) =>
|
||||
b.textContent?.includes(`${K}.title`),
|
||||
)!;
|
||||
const expand = async () => {
|
||||
if (disclosure().getAttribute('aria-expanded') === 'false')
|
||||
await act(async () => disclosure().click());
|
||||
};
|
||||
const groups = () =>
|
||||
Array.from(container.querySelectorAll<HTMLElement>('[data-tool]'));
|
||||
const groupTitles = () =>
|
||||
groups().map((g) => g.querySelector('h4')?.textContent);
|
||||
const group = (id: string) =>
|
||||
container.querySelector<HTMLElement>(`[data-tool="${id}"]`)!;
|
||||
const choice = (id: string, value: 'off' | 'all') =>
|
||||
group(id).querySelector<HTMLButtonElement>(`[data-choice="${value}"]`)!;
|
||||
const customize = async (id: string) => {
|
||||
const button = Array.from(group(id).querySelectorAll('button')).find(
|
||||
(b) => b.getAttribute('aria-expanded') === 'false',
|
||||
);
|
||||
if (button) await act(async () => button.click());
|
||||
};
|
||||
const actionSwitch = (id: string, label: string) => {
|
||||
const row = Array.from(
|
||||
group(id).querySelectorAll<HTMLElement>('[data-slot="setting-row"]'),
|
||||
).find((r) => r.querySelector('label')?.textContent === label)!;
|
||||
return row.querySelector<HTMLButtonElement>('[role="switch"]')!;
|
||||
};
|
||||
const summary = () =>
|
||||
container.querySelector('[data-testid="api-writes-summary"]')?.textContent;
|
||||
const savedConfig = () =>
|
||||
JSON.parse(
|
||||
(updateAgent.mock.calls.at(-1)![1] as FormData).get('config') as string,
|
||||
);
|
||||
|
||||
it('starts folded to a summary line, and nothing is allowed by default', async () => {
|
||||
await render(agent());
|
||||
expect(disclosure().getAttribute('aria-expanded')).toBe('false');
|
||||
expect(summary()).toBe(`${K}.summaryNone`);
|
||||
expect(groups()).toHaveLength(0);
|
||||
await expand();
|
||||
expect(disclosure().getAttribute('aria-expanded')).toBe('true');
|
||||
expect(choice('tg', 'off').getAttribute('data-state')).toBe('on');
|
||||
});
|
||||
|
||||
it('opens straight away when asked to', async () => {
|
||||
await render(agent(), { defaultOpen: true });
|
||||
expect(disclosure().getAttribute('aria-expanded')).toBe('true');
|
||||
expect(groupTitles()).toEqual(['Telegram']);
|
||||
});
|
||||
|
||||
it('names the tools that can make changes and counts what is allowed', async () => {
|
||||
await render(
|
||||
agent({
|
||||
tools: ['tg', 'crm', 'memory'],
|
||||
config: {
|
||||
api_write_allowlist: ['tg:telegram_send_message', 'crm:create_lead'],
|
||||
},
|
||||
} as Partial<Agent>),
|
||||
);
|
||||
expect(summary()).toBe(
|
||||
`${K}.summaryTools(tools=Telegram and CRM API) · ${K}.summaryCount(allowed=2,formatted=3)`,
|
||||
);
|
||||
});
|
||||
|
||||
it("groups only writes on the owner's credentials by tool, each action under Customize", async () => {
|
||||
await render(agent());
|
||||
await expand();
|
||||
expect(groupTitles()).toEqual(['Telegram']);
|
||||
expect(container.textContent).not.toContain('Telegram send message');
|
||||
await customize('tg');
|
||||
const labels = Array.from(
|
||||
group('tg').querySelectorAll('[data-slot="setting-row"] label'),
|
||||
).map((l) => l.textContent);
|
||||
expect(labels).toEqual(['Telegram send message', 'Telegram pin message']);
|
||||
expect(group('tg').textContent).toContain('Sends a message.');
|
||||
expect(
|
||||
actionSwitch('tg', 'Telegram send message').getAttribute('aria-checked'),
|
||||
).toBe('false');
|
||||
});
|
||||
|
||||
it("allows all of a tool's changes at once, without dropping the rest of the config", async () => {
|
||||
updateAgent.mockResolvedValue({ ok: true });
|
||||
const onConfigChange = await render(agent());
|
||||
await expand();
|
||||
await act(async () => choice('tg', 'all').click());
|
||||
const expected = {
|
||||
guardrails: { controls: [] },
|
||||
api_write_allowlist: [
|
||||
'tg:telegram_send_message',
|
||||
'tg:telegram_pin_message',
|
||||
],
|
||||
};
|
||||
expect(updateAgent).toHaveBeenCalledTimes(1);
|
||||
expect(savedConfig()).toEqual(expected);
|
||||
expect(onConfigChange).toHaveBeenCalledWith(expected);
|
||||
expect(choice('tg', 'all').getAttribute('data-state')).toBe('on');
|
||||
expect(summary()).toBe(
|
||||
`${K}.summaryTools(tools=Telegram) · ${K}.summaryCount(allowed=2,formatted=2)`,
|
||||
);
|
||||
});
|
||||
|
||||
it('locks the choices while a save is in flight, so saves never overlap', async () => {
|
||||
let finish: (value: { ok: boolean }) => void = () => undefined;
|
||||
updateAgent.mockReturnValue(
|
||||
new Promise((resolve) => {
|
||||
finish = resolve;
|
||||
}),
|
||||
);
|
||||
await render(agent());
|
||||
await expand();
|
||||
await act(async () => choice('tg', 'all').click());
|
||||
expect(choice('tg', 'off').hasAttribute('disabled')).toBe(true);
|
||||
await act(async () => choice('tg', 'off').click());
|
||||
expect(updateAgent).toHaveBeenCalledTimes(1);
|
||||
await act(async () => finish({ ok: true }));
|
||||
expect(choice('tg', 'off').hasAttribute('disabled')).toBe(false);
|
||||
});
|
||||
|
||||
it('turns a tool off again, keeping entries of tools no longer listed', async () => {
|
||||
updateAgent.mockResolvedValue({ ok: true });
|
||||
await render(
|
||||
agent({
|
||||
config: {
|
||||
api_write_allowlist: [
|
||||
'tg:telegram_send_message',
|
||||
'tg:telegram_pin_message',
|
||||
'old:gone_action',
|
||||
],
|
||||
},
|
||||
} as Partial<Agent>),
|
||||
);
|
||||
await expand();
|
||||
await act(async () => choice('tg', 'off').click());
|
||||
expect(savedConfig().api_write_allowlist).toEqual(['old:gone_action']);
|
||||
});
|
||||
|
||||
// Every tool has a Customize link; each names its tool to screen readers.
|
||||
it('ties each Customize link to its tool', async () => {
|
||||
await render(agent({ tools: ['tg', 'crm'] }), { defaultOpen: true });
|
||||
for (const [id, name] of [
|
||||
['tg', 'Telegram'],
|
||||
['crm', 'CRM API'],
|
||||
]) {
|
||||
const link = Array.from(group(id).querySelectorAll('button')).find(
|
||||
(b) => b.getAttribute('aria-expanded') === 'false',
|
||||
)!;
|
||||
const describedBy = link.getAttribute('aria-describedby');
|
||||
expect(describedBy).toBeTruthy();
|
||||
expect(document.getElementById(describedBy!)?.textContent).toBe(name);
|
||||
}
|
||||
});
|
||||
|
||||
it('allows one action under Customize, leaving the tool choice mixed', async () => {
|
||||
updateAgent.mockResolvedValue({ ok: true });
|
||||
await render(agent());
|
||||
await expand();
|
||||
await customize('tg');
|
||||
await act(async () => actionSwitch('tg', 'Telegram pin message').click());
|
||||
expect(savedConfig().api_write_allowlist).toEqual([
|
||||
'tg:telegram_pin_message',
|
||||
]);
|
||||
expect(
|
||||
actionSwitch('tg', 'Telegram pin message').getAttribute('aria-checked'),
|
||||
).toBe('true');
|
||||
expect(choice('tg', 'off').getAttribute('data-state')).toBe('off');
|
||||
expect(choice('tg', 'all').getAttribute('data-state')).toBe('off');
|
||||
});
|
||||
|
||||
it('saves on top of the last saved config, not unsaved form edits', async () => {
|
||||
updateAgent.mockResolvedValue({ ok: true });
|
||||
const saved = { guardrails: { controls: [] } } as unknown as NonNullable<
|
||||
Agent['config']
|
||||
>;
|
||||
const draft = agent({
|
||||
config: { guardrails: { controls: [{ id: 'unsaved' }] } },
|
||||
} as unknown as Partial<Agent>);
|
||||
const onConfigChange = await render(draft, {
|
||||
getSavedConfig: () => saved,
|
||||
});
|
||||
await expand();
|
||||
await act(async () => choice('tg', 'all').click());
|
||||
const expected = {
|
||||
guardrails: { controls: [] },
|
||||
api_write_allowlist: [
|
||||
'tg:telegram_send_message',
|
||||
'tg:telegram_pin_message',
|
||||
],
|
||||
};
|
||||
expect(savedConfig()).toEqual(expected);
|
||||
expect(onConfigChange).toHaveBeenCalledWith(expected);
|
||||
});
|
||||
|
||||
it('puts the choice back and says so when saving fails', async () => {
|
||||
updateAgent.mockResolvedValue({ ok: false });
|
||||
await render(agent());
|
||||
await expand();
|
||||
await act(async () => choice('tg', 'all').click());
|
||||
expect(choice('tg', 'off').getAttribute('data-state')).toBe('on');
|
||||
expect(summary()).toBe(`${K}.summaryNone`);
|
||||
expect(selectActionToast(store.getState())?.variant).toBe('destructive');
|
||||
});
|
||||
|
||||
it('lists writes on stored credentials of tools without a connection', async () => {
|
||||
await render(agent({ tools: ['crm'] }), { defaultOpen: true });
|
||||
expect(groupTitles()).toEqual(['CRM API']);
|
||||
await customize('crm');
|
||||
expect(
|
||||
group('crm').querySelector('[data-slot="setting-row"] label')
|
||||
?.textContent,
|
||||
).toBe('Create lead');
|
||||
});
|
||||
|
||||
it("lists a sponsor's tool the owner can't see, but not a stopped one", async () => {
|
||||
await render(agent({ tools: ['bob-jira', 'gone'] }), {
|
||||
defaultOpen: true,
|
||||
});
|
||||
expect(groupTitles()).toEqual(['Bob Jira']);
|
||||
});
|
||||
|
||||
it("lists writes of a workflow agent's node tools", async () => {
|
||||
getAgent.mockResolvedValue(
|
||||
readAgent([], { agent_type: 'workflow', workflow: 'w1' }),
|
||||
);
|
||||
getWorkflow.mockResolvedValue({
|
||||
ok: true,
|
||||
json: async () => ({
|
||||
data: {
|
||||
resource_states: [
|
||||
state({
|
||||
id: 'node-tool',
|
||||
name: 'Node Slack',
|
||||
owner_credential_writes: ['post_message'],
|
||||
}),
|
||||
],
|
||||
},
|
||||
}),
|
||||
});
|
||||
await render(agent({ tools: [] }), { defaultOpen: true });
|
||||
expect(groupTitles()).toEqual(['Node Slack']);
|
||||
});
|
||||
|
||||
it('renders nothing for an agent without connected tools', async () => {
|
||||
await render(agent({ tools: ['memory'] }));
|
||||
expect(container.textContent).toBe('');
|
||||
});
|
||||
});
|
||||
Loaded 100 of 317 files, more files were not shown because too many files have changed in this diff.
Show more
Reference in new issue
Block a user