mirror of
https://github.com/Tencent/WeKnora.git
synced 2026-10-04 06:48:15 +08:00
feat(plugin): remote runtime, Python SDK and python host plugins (P2 batch 3) (#3724)
* feat(plugin): remote runtime * feat(plugin): Python SDK and python host plugins * fix(pluginsdk): skip reverse DNS when a Python plugin binds
This commit is contained in:
@@ -1,17 +1,21 @@
|
||||
name: Plugin SDK
|
||||
|
||||
# Builds and tests the pluginsdk/ Go module (extension protocol, Go SDK,
|
||||
# client and conformance suite) on every push / PR that touches it.
|
||||
# client and conformance suite) and the Python SDK in pluginsdk/python on
|
||||
# every push / PR that touches them. The Go conformance suite also runs the
|
||||
# Python SDK's test plugin.
|
||||
|
||||
on:
|
||||
push:
|
||||
branches: [main]
|
||||
paths:
|
||||
- 'pluginsdk/**'
|
||||
- 'examples/plugins/notebooks/**'
|
||||
- '.github/workflows/plugin-sdk.yml'
|
||||
pull_request:
|
||||
paths:
|
||||
- 'pluginsdk/**'
|
||||
- 'examples/plugins/notebooks/**'
|
||||
- '.github/workflows/plugin-sdk.yml'
|
||||
|
||||
defaults:
|
||||
@@ -36,9 +40,38 @@ jobs:
|
||||
with:
|
||||
go-version: '1.26'
|
||||
cache-dependency-path: pluginsdk/go.mod
|
||||
- uses: actions/setup-python@v6
|
||||
with:
|
||||
python-version: '3.13'
|
||||
- name: go build
|
||||
run: go build ./...
|
||||
- name: go test
|
||||
run: go test -race ./...
|
||||
- name: go vet
|
||||
run: go vet ./...
|
||||
|
||||
python:
|
||||
name: python ${{ matrix.python }} (${{ matrix.os }})
|
||||
runs-on: ${{ matrix.os }}
|
||||
strategy:
|
||||
fail-fast: false
|
||||
matrix:
|
||||
include:
|
||||
- { os: ubuntu-latest, python: '3.9' }
|
||||
- { os: ubuntu-latest, python: '3.13' }
|
||||
- { os: windows-latest, python: '3.13' }
|
||||
defaults:
|
||||
run:
|
||||
working-directory: pluginsdk/python
|
||||
steps:
|
||||
- uses: actions/checkout@v6
|
||||
- uses: actions/setup-python@v6
|
||||
with:
|
||||
python-version: ${{ matrix.python }}
|
||||
- name: unittest
|
||||
run: python -W error::ResourceWarning -m unittest discover -s tests -v
|
||||
- name: example plugin
|
||||
working-directory: examples/plugins/notebooks
|
||||
run: python -m unittest -v test_main
|
||||
- name: build the package
|
||||
run: python -m pip install build && python -m build
|
||||
|
||||
@@ -0,0 +1,106 @@
|
||||
"""A WeKnora parser plugin, written with the Python SDK: Jupyter notebooks
|
||||
(.ipynb) become Markdown. Markdown cells are kept as they are, code cells
|
||||
become fenced blocks, text outputs follow their cell and PNG/JPEG charts are
|
||||
handed to WeKnora as images."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import json
|
||||
|
||||
from weknora_plugin import ErrorCode, ParsedImage, ParseOutput, Plugin, PluginError
|
||||
|
||||
plugin = Plugin("weknora-examples.notebooks", "1.0.0")
|
||||
|
||||
IMAGE_TYPES = {"image/png": "png", "image/jpeg": "jpg"}
|
||||
|
||||
|
||||
def _text(value) -> str:
|
||||
"""Notebook strings are either a string or a list of lines."""
|
||||
if isinstance(value, list):
|
||||
return "".join(value)
|
||||
return value or ""
|
||||
|
||||
|
||||
def _fence(body: str, lang: str = "") -> str:
|
||||
# A fence longer than any backtick run inside, so code cannot close it.
|
||||
longest, run = 0, 0
|
||||
for ch in body:
|
||||
run = run + 1 if ch == "`" else 0
|
||||
longest = max(longest, run)
|
||||
ticks = "`" * max(3, longest + 1)
|
||||
return f"{ticks}{lang}\n{body.rstrip()}\n{ticks}"
|
||||
|
||||
|
||||
def convert(nb: dict, include_outputs: bool = True, max_output_chars: int = 2000) -> ParseOutput:
|
||||
lang = ((nb.get("metadata") or {}).get("kernelspec") or {}).get("language") or (
|
||||
(nb.get("metadata") or {}).get("language_info") or {}
|
||||
).get("name", "")
|
||||
parts, images = [], []
|
||||
for i, cell in enumerate(nb.get("cells") or []):
|
||||
kind, source = cell.get("cell_type"), _text(cell.get("source")).strip()
|
||||
if kind == "markdown" and source:
|
||||
parts.append(source)
|
||||
elif kind == "code":
|
||||
if source:
|
||||
parts.append(_fence(source, lang))
|
||||
if include_outputs:
|
||||
parts.extend(_outputs(i, cell.get("outputs") or [], images, max_output_chars))
|
||||
elif kind == "raw" and source:
|
||||
parts.append(_fence(source))
|
||||
title = next((p.splitlines()[0][2:].strip() for p in parts if p.startswith("# ")), "")
|
||||
meta = {"cells": str(len(nb.get("cells") or [])), "language": lang}
|
||||
if title:
|
||||
meta["title"] = title
|
||||
return ParseOutput(markdown="\n\n".join(parts) + "\n", images=images, metadata=meta)
|
||||
|
||||
|
||||
def _outputs(cell: int, outputs: list, images: list, limit: int) -> list:
|
||||
parts = []
|
||||
for j, out in enumerate(outputs):
|
||||
kind = out.get("output_type")
|
||||
if kind == "stream":
|
||||
text = _text(out.get("text"))
|
||||
elif kind == "error":
|
||||
text = f"{out.get('ename', 'Error')}: {out.get('evalue', '')}"
|
||||
else:
|
||||
data = out.get("data") or {}
|
||||
image = next((m for m in IMAGE_TYPES if m in data), None)
|
||||
if image:
|
||||
ref = f"images/cell{cell}-{j}.{IMAGE_TYPES[image]}"
|
||||
try:
|
||||
images.append(ParsedImage(original_ref=ref, data=base64.b64decode(_text(data[image])), mime_type=image))
|
||||
parts.append(f"")
|
||||
except ValueError:
|
||||
pass
|
||||
continue
|
||||
if "text/markdown" in data:
|
||||
parts.append(_text(data["text/markdown"]).strip())
|
||||
continue
|
||||
text = _text(data.get("text/plain"))
|
||||
text = text.strip()
|
||||
if text:
|
||||
if limit and len(text) > limit:
|
||||
text = text[:limit] + "\n…"
|
||||
parts.append(_fence(text, "text"))
|
||||
return parts
|
||||
|
||||
|
||||
@plugin.parser("ipynb")
|
||||
def parse(call, doc):
|
||||
try:
|
||||
nb = json.loads(doc.content.decode("utf-8-sig"))
|
||||
except (UnicodeDecodeError, ValueError) as e:
|
||||
raise PluginError(ErrorCode.BAD_REQUEST, f"{doc.file_name or 'file'} is not a Jupyter notebook: {e}") from None
|
||||
if not isinstance(nb, dict) or "cells" not in nb:
|
||||
raise PluginError(ErrorCode.BAD_REQUEST, f"{doc.file_name or 'file'} is not a Jupyter notebook (no cells)")
|
||||
cfg = call.tenant
|
||||
return convert(
|
||||
nb,
|
||||
include_outputs=cfg.get("include_outputs", True) is not False,
|
||||
max_output_chars=int(cfg.get("max_output_chars", 2000) or 0),
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
plugin.serve()
|
||||
Executable
+18
@@ -0,0 +1,18 @@
|
||||
#!/usr/bin/env bash
|
||||
# Packages the notebooks example plugin as a .wkp: plugin.yaml, schemas,
|
||||
# main.py and the Python SDK vendored under vendor/. Nothing is compiled;
|
||||
# WeKnora runs main.py with its own python3.
|
||||
set -euo pipefail
|
||||
cd "$(dirname "$0")"
|
||||
version=$(sed -n 's/^version: //p' plugin.yaml)
|
||||
out=${1:-"weknora-examples-notebooks-${version}.wkp"}
|
||||
stage=$(mktemp -d)
|
||||
trap 'rm -rf "$stage"' EXIT
|
||||
cp plugin.yaml main.py "$stage"/
|
||||
cp -r schemas "$stage"/
|
||||
mkdir -p "$stage/vendor"
|
||||
cp -r ../../../pluginsdk/python/src/weknora_plugin "$stage/vendor/"
|
||||
find "$stage" -name __pycache__ -prune -exec rm -rf {} +
|
||||
rm -f "$out"
|
||||
(cd "$stage" && zip -qr - .) > "$out"
|
||||
echo "wrote $out"
|
||||
@@ -0,0 +1,22 @@
|
||||
schemaVersion: 1
|
||||
id: weknora-examples.notebooks
|
||||
version: 1.0.0
|
||||
apiVersion: weknora.plugin/v1
|
||||
name: { en-US: Jupyter notebooks, zh-CN: Jupyter 笔记本 }
|
||||
description:
|
||||
en-US: Parse Jupyter notebooks into Markdown with their code, text outputs and charts. Written in Python.
|
||||
zh-CN: 把 Jupyter 笔记本解析为 Markdown,保留代码、文本输出和图表。用 Python 编写。
|
||||
publisher: { id: weknora-examples, name: WeKnora examples }
|
||||
homepage: https://github.com/Tencent/WeKnora/tree/main/examples/plugins/notebooks
|
||||
license: MIT
|
||||
runtime: { type: host, kind: python, entry: main.py }
|
||||
config:
|
||||
tenant: schemas/tenant.yaml
|
||||
contributes:
|
||||
parsers:
|
||||
- id: ipynb
|
||||
name: { en-US: Jupyter notebooks, zh-CN: Jupyter 笔记本 }
|
||||
description:
|
||||
en-US: Markdown and code cells, text outputs and PNG/JPEG charts
|
||||
zh-CN: Markdown 与代码单元格、文本输出和 PNG/JPEG 图表
|
||||
fileTypes: [ipynb]
|
||||
@@ -0,0 +1,19 @@
|
||||
type: object
|
||||
properties:
|
||||
include_outputs:
|
||||
type: boolean
|
||||
title: Include cell outputs
|
||||
description: Index what code cells printed and plotted, not just the code
|
||||
default: true
|
||||
x-i18n:
|
||||
title: { zh-CN: 包含单元格输出 }
|
||||
description: { zh-CN: 除代码外,也索引代码单元格的打印内容和图表 }
|
||||
max_output_chars:
|
||||
type: integer
|
||||
title: Output length limit
|
||||
description: Characters kept from each text output
|
||||
default: 2000
|
||||
minimum: 0
|
||||
x-i18n:
|
||||
title: { zh-CN: 输出长度上限 }
|
||||
description: { zh-CN: 每段文本输出保留的字符数 }
|
||||
@@ -0,0 +1,73 @@
|
||||
import base64
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
import unittest
|
||||
|
||||
HERE = os.path.dirname(os.path.abspath(__file__))
|
||||
sys.path.insert(0, os.path.join(HERE, "..", "..", "..", "pluginsdk", "python", "src"))
|
||||
sys.path.insert(0, HERE)
|
||||
|
||||
from main import convert, plugin # noqa: E402
|
||||
from weknora_plugin import ErrorCode, ParseInput, PluginError # noqa: E402
|
||||
from weknora_plugin.plugin import Call # noqa: E402
|
||||
|
||||
PNG = base64.b64encode(b"\x89PNG fake").decode()
|
||||
|
||||
NOTEBOOK = {
|
||||
"metadata": {"kernelspec": {"language": "python"}},
|
||||
"cells": [
|
||||
{"cell_type": "markdown", "source": ["# Sales report\n", "Quarterly numbers."]},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"source": "print('total', 42)\nx = '```'",
|
||||
"outputs": [
|
||||
{"output_type": "stream", "name": "stdout", "text": ["total 42\n"]},
|
||||
{"output_type": "display_data", "data": {"image/png": PNG, "text/plain": "<Figure>"}},
|
||||
{"output_type": "execute_result", "data": {"text/plain": ["'" + "a" * 50 + "'"]}},
|
||||
{"output_type": "error", "ename": "ValueError", "evalue": "bad"},
|
||||
],
|
||||
},
|
||||
{"cell_type": "code", "source": "", "outputs": []},
|
||||
],
|
||||
}
|
||||
|
||||
|
||||
class ConvertTest(unittest.TestCase):
|
||||
def test_cells_and_outputs(self):
|
||||
out = convert(NOTEBOOK, max_output_chars=20)
|
||||
md = out.markdown
|
||||
self.assertTrue(md.startswith("# Sales report\nQuarterly numbers.\n\n````python\nprint('total', 42)"))
|
||||
self.assertIn("```text\ntotal 42\n```", md)
|
||||
self.assertIn("", md)
|
||||
self.assertIn("'" + "a" * 19 + "\n…", md)
|
||||
self.assertIn("ValueError: bad", md)
|
||||
self.assertEqual([(i.original_ref, i.data, i.mime_type) for i in out.images], [("images/cell1-1.png", b"\x89PNG fake", "image/png")])
|
||||
self.assertEqual(out.metadata, {"cells": "3", "language": "python", "title": "Sales report"})
|
||||
|
||||
def test_without_outputs(self):
|
||||
out = convert(NOTEBOOK, include_outputs=False)
|
||||
self.assertNotIn("total 42", out.markdown)
|
||||
self.assertEqual(out.images, [])
|
||||
|
||||
|
||||
class ParseTest(unittest.TestCase):
|
||||
def call(self, tenant=None):
|
||||
return Call({"context": {"tenantId": 1}, "config": {"tenant": tenant or {}}})
|
||||
|
||||
def test_parse(self):
|
||||
doc = ParseInput(file_name="r.ipynb", file_type="ipynb", content=json.dumps(NOTEBOOK).encode())
|
||||
parse = plugin._parsers["ipynb"]
|
||||
self.assertIn("total 42", parse(self.call(), doc).markdown)
|
||||
self.assertNotIn("total 42", parse(self.call({"include_outputs": False}), doc).markdown)
|
||||
|
||||
def test_not_a_notebook(self):
|
||||
parse = plugin._parsers["ipynb"]
|
||||
for content in (b"WeKnora conformance check\n", b"[]", b"\xff\xfe"):
|
||||
with self.assertRaises(PluginError) as ctx:
|
||||
parse(self.call(), ParseInput(file_name="x.ipynb", content=content))
|
||||
self.assertEqual(ctx.exception.code, ErrorCode.BAD_REQUEST)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -51,15 +51,20 @@ export interface InstalledPlugin {
|
||||
manifest?: PluginManifest
|
||||
versions: PluginVersion[]
|
||||
node?: PluginNodeStatus
|
||||
/** Where a remote plugin's service runs. */
|
||||
remote_url?: string
|
||||
/** A remote plugin's new signing secret, only in the response that issued it. */
|
||||
issuedSecret?: string
|
||||
}
|
||||
|
||||
/** A package to inspect or install: an uploaded file or a URL. */
|
||||
export type PackageSource = { file: File } | { url: string }
|
||||
|
||||
function packageForm(file: File, digest?: string) {
|
||||
function packageForm(file: File, digest?: string, remoteUrl?: string) {
|
||||
const form = new FormData()
|
||||
form.append('file', file)
|
||||
if (digest) form.append('digest', digest)
|
||||
if (remoteUrl) form.append('remote_url', remoteUrl)
|
||||
return form
|
||||
}
|
||||
|
||||
@@ -81,14 +86,21 @@ export function inspectPluginPackage(source: PackageSource) {
|
||||
return post<{ data: PluginPreview }>(`${BASE}/inspect`, { url: source.url }, { timeout: PACKAGE_TIMEOUT })
|
||||
}
|
||||
|
||||
/** Installs the package reviewed with inspect; digest pins it to that package. */
|
||||
export function installPluginPackage(source: PackageSource, digest: string) {
|
||||
/**
|
||||
* Installs the package reviewed with inspect; digest pins it to that package.
|
||||
* A remote plugin also needs the URL of its service (optional on upgrades).
|
||||
*/
|
||||
export function installPluginPackage(source: PackageSource, digest: string, remoteUrl?: string) {
|
||||
if ('file' in source) {
|
||||
return postUpload(BASE, packageForm(source.file, digest), undefined, { timeout: PACKAGE_TIMEOUT }) as Promise<{
|
||||
data: InstalledPlugin
|
||||
}>
|
||||
return postUpload(BASE, packageForm(source.file, digest, remoteUrl), undefined, {
|
||||
timeout: PACKAGE_TIMEOUT,
|
||||
}) as Promise<{ data: InstalledPlugin }>
|
||||
}
|
||||
return post<{ data: InstalledPlugin }>(BASE, { url: source.url, digest }, { timeout: PACKAGE_TIMEOUT })
|
||||
return post<{ data: InstalledPlugin }>(
|
||||
BASE,
|
||||
{ url: source.url, digest, remote_url: remoteUrl || undefined },
|
||||
{ timeout: PACKAGE_TIMEOUT },
|
||||
)
|
||||
}
|
||||
|
||||
export function setInstalledPluginEnabled(id: string, enabled: boolean) {
|
||||
@@ -99,6 +111,15 @@ export function activatePluginVersion(id: string, version: string) {
|
||||
return put<{ data: InstalledPlugin }>(`${BASE}/${encodeURIComponent(id)}/active-version`, { version })
|
||||
}
|
||||
|
||||
export function setPluginRemoteUrl(id: string, url: string) {
|
||||
return put<{ data: InstalledPlugin }>(`${BASE}/${encodeURIComponent(id)}/remote-url`, { url })
|
||||
}
|
||||
|
||||
/** Issues a new signing secret; the response carries it in issuedSecret. */
|
||||
export function rotatePluginSecret(id: string) {
|
||||
return post<{ data: InstalledPlugin }>(`${BASE}/${encodeURIComponent(id)}/secret/rotate`, {})
|
||||
}
|
||||
|
||||
export function uninstallPlugin(id: string) {
|
||||
return del(`${BASE}/${encodeURIComponent(id)}`)
|
||||
}
|
||||
|
||||
@@ -2736,7 +2736,8 @@ export default {
|
||||
pluginAdmin: {
|
||||
runtime: {
|
||||
declarative: 'Declarative',
|
||||
host: 'Local process'
|
||||
host: 'Local process',
|
||||
remote: 'Remote service'
|
||||
},
|
||||
title: 'Plugin management',
|
||||
description: 'Install and manage plugins for the whole platform. Every workspace sees an installed plugin, and each one enables it for itself; disabling a plugin here unloads it on every node.',
|
||||
@@ -2780,6 +2781,13 @@ export default {
|
||||
hostApi: 'WeKnora API',
|
||||
events: 'Events'
|
||||
},
|
||||
secret: {
|
||||
title: 'Plugin signing secret',
|
||||
description: 'Set this secret as the plugin service\'s WEKNORA_PLUGIN_SECRET environment variable; the service uses it to verify that requests come from WeKnora. It is shown only once and cannot be viewed again after you close this.',
|
||||
copy: 'Copy',
|
||||
copied: 'Copied',
|
||||
done: 'I have saved it'
|
||||
},
|
||||
install: {
|
||||
title: 'Install plugin',
|
||||
description: 'Upload a .wkp package or give its download URL, review it, then install.',
|
||||
@@ -2794,6 +2802,9 @@ export default {
|
||||
fileHint: 'A .wkp file (a zip with plugin.yaml), up to 64 MB.',
|
||||
urlLabel: 'Download URL',
|
||||
urlHint: 'The server downloads this URL; private network addresses are refused.',
|
||||
remoteUrlLabel: 'Service URL',
|
||||
remoteUrlHint: 'The plugin service\'s HTTP(S) address. Private network hosts must be listed in SSRF_WHITELIST.',
|
||||
remoteUrlKeep: 'On an upgrade, leave empty to keep the current URL',
|
||||
reviewSection: 'Review before installing',
|
||||
noPermissions: 'This plugin asks for no extra permissions.',
|
||||
configNotice: 'This plugin needs configuration: platform settings are in its details after installing, and workspace admins fill in workspace settings in Plugins.',
|
||||
@@ -2836,7 +2847,16 @@ export default {
|
||||
uninstall: 'Uninstall plugin',
|
||||
uninstallConfirm: 'Uninstall? Every node unloads the plugin right away.',
|
||||
uninstalled: 'Plugin uninstalled',
|
||||
uninstallFailed: 'Uninstall failed'
|
||||
uninstallFailed: 'Uninstall failed',
|
||||
remote: 'Remote service',
|
||||
remoteUrl: 'Service URL',
|
||||
editUrl: 'Change',
|
||||
urlSaved: 'Service URL updated',
|
||||
urlSaveFailed: 'Failed to update the service URL',
|
||||
rotateSecret: 'Rotate secret',
|
||||
rotateHint: 'Issues a new signing secret. Calls to the service fail until it is given the new one.',
|
||||
rotateConfirm: 'Rotate the secret? The old one stops working immediately.',
|
||||
rotateFailed: 'Failed to rotate the secret'
|
||||
}
|
||||
},
|
||||
pluginCenter: {
|
||||
|
||||
@@ -2736,7 +2736,8 @@ export default {
|
||||
pluginAdmin: {
|
||||
runtime: {
|
||||
declarative: '宣言型',
|
||||
host: 'ローカルプロセス'
|
||||
host: 'ローカルプロセス',
|
||||
remote: 'リモートサービス'
|
||||
},
|
||||
title: 'プラグイン管理',
|
||||
description: 'プラットフォーム全体のプラグインをインストール・管理します。インストールしたプラグインはすべてのワークスペースに表示され、各ワークスペースで個別に有効化します。ここで無効にすると全ノードでアンロードされます。',
|
||||
@@ -2780,6 +2781,13 @@ export default {
|
||||
hostApi: 'WeKnora API',
|
||||
events: 'イベント'
|
||||
},
|
||||
secret: {
|
||||
title: 'プラグイン署名シークレット',
|
||||
description: 'このシークレットをプラグインサービスの環境変数 WEKNORA_PLUGIN_SECRET に設定してください。サービスはこれでリクエストが WeKnora からのものか確認します。表示されるのは今回だけで、閉じると再表示できません。',
|
||||
copy: 'コピー',
|
||||
copied: 'コピーしました',
|
||||
done: '保存しました'
|
||||
},
|
||||
install: {
|
||||
title: 'プラグインをインストール',
|
||||
description: '.wkp パッケージをアップロードするかダウンロード URL を指定し、確認してからインストールします。',
|
||||
@@ -2794,6 +2802,9 @@ export default {
|
||||
fileHint: '.wkp ファイル(plugin.yaml を含む zip)、最大 64 MB。',
|
||||
urlLabel: 'ダウンロード URL',
|
||||
urlHint: 'サーバーがこの URL からダウンロードします。プライベートネットワークのアドレスは拒否されます。',
|
||||
remoteUrlLabel: 'サービス URL',
|
||||
remoteUrlHint: 'プラグインサービスの HTTP(S) アドレス。プライベートネットワークのホストは SSRF_WHITELIST に追加する必要があります。',
|
||||
remoteUrlKeep: 'アップグレード時は空欄のままで現在の URL を引き続き使用します',
|
||||
reviewSection: 'インストール前の確認',
|
||||
noPermissions: 'このプラグインは追加の権限を要求しません。',
|
||||
configNotice: 'このプラグインには設定が必要です。プラットフォーム設定はインストール後の詳細で、ワークスペース設定は各ワークスペース管理者がプラグインセンターで入力します。',
|
||||
@@ -2836,7 +2847,16 @@ export default {
|
||||
uninstall: 'プラグインをアンインストール',
|
||||
uninstallConfirm: 'アンインストールしますか?すべてのノードで直ちにアンロードされます。',
|
||||
uninstalled: 'プラグインをアンインストールしました',
|
||||
uninstallFailed: 'アンインストールに失敗しました'
|
||||
uninstallFailed: 'アンインストールに失敗しました',
|
||||
remote: 'リモートサービス',
|
||||
remoteUrl: 'サービス URL',
|
||||
editUrl: '変更',
|
||||
urlSaved: 'サービス URL を更新しました',
|
||||
urlSaveFailed: 'サービス URL の更新に失敗しました',
|
||||
rotateSecret: 'シークレットをローテーション',
|
||||
rotateHint: '新しい署名シークレットを発行します。サービスに新しいシークレットを設定するまで呼び出しは失敗します。',
|
||||
rotateConfirm: 'ローテーションしますか?古いシークレットは直ちに無効になります。',
|
||||
rotateFailed: 'シークレットのローテーションに失敗しました'
|
||||
}
|
||||
},
|
||||
pluginCenter: {
|
||||
|
||||
@@ -5254,7 +5254,8 @@ export default {
|
||||
pluginAdmin: {
|
||||
runtime: {
|
||||
declarative: '선언형',
|
||||
host: '로컬 프로세스'
|
||||
host: '로컬 프로세스',
|
||||
remote: '원격 서비스'
|
||||
},
|
||||
title: '플러그인 관리',
|
||||
description: '플랫폼 전체의 플러그인을 설치하고 관리합니다. 설치된 플러그인은 모든 워크스페이스에 보이며 각 워크스페이스가 직접 활성화합니다. 여기서 비활성화하면 모든 노드에서 언로드됩니다.',
|
||||
@@ -5298,6 +5299,13 @@ export default {
|
||||
hostApi: 'WeKnora API',
|
||||
events: '이벤트'
|
||||
},
|
||||
secret: {
|
||||
title: '플러그인 서명 비밀 키',
|
||||
description: '이 비밀 키를 플러그인 서비스의 환경 변수 WEKNORA_PLUGIN_SECRET으로 설정하세요. 서비스는 이를 사용해 요청이 WeKnora에서 온 것인지 확인합니다. 이번에만 표시되며 닫은 후에는 다시 볼 수 없습니다.',
|
||||
copy: '복사',
|
||||
copied: '복사했습니다',
|
||||
done: '저장했습니다'
|
||||
},
|
||||
install: {
|
||||
title: '플러그인 설치',
|
||||
description: '.wkp 패키지를 업로드하거나 다운로드 URL을 입력하고, 검토한 뒤 설치합니다.',
|
||||
@@ -5312,6 +5320,9 @@ export default {
|
||||
fileHint: '.wkp 파일(plugin.yaml이 포함된 zip), 최대 64 MB.',
|
||||
urlLabel: '다운로드 URL',
|
||||
urlHint: '서버가 이 URL에서 다운로드합니다. 사설 네트워크 주소는 거부됩니다.',
|
||||
remoteUrlLabel: '서비스 URL',
|
||||
remoteUrlHint: '플러그인 서비스의 HTTP(S) 주소입니다. 사설 네트워크 호스트는 SSRF_WHITELIST에 추가해야 합니다.',
|
||||
remoteUrlKeep: '업그레이드 시 비워 두면 현재 URL을 계속 사용합니다',
|
||||
reviewSection: '설치 전 검토',
|
||||
noPermissions: '이 플러그인은 추가 권한을 요청하지 않습니다.',
|
||||
configNotice: '이 플러그인은 설정이 필요합니다. 플랫폼 설정은 설치 후 상세 화면에서, 워크스페이스 설정은 각 워크스페이스 관리자가 플러그인 센터에서 입력합니다.',
|
||||
@@ -5354,7 +5365,16 @@ export default {
|
||||
uninstall: '플러그인 제거',
|
||||
uninstallConfirm: '제거할까요? 모든 노드에서 즉시 언로드됩니다.',
|
||||
uninstalled: '플러그인을 제거했습니다',
|
||||
uninstallFailed: '제거하지 못했습니다'
|
||||
uninstallFailed: '제거하지 못했습니다',
|
||||
remote: '원격 서비스',
|
||||
remoteUrl: '서비스 URL',
|
||||
editUrl: '변경',
|
||||
urlSaved: '서비스 URL을 업데이트했습니다',
|
||||
urlSaveFailed: '서비스 URL을 업데이트하지 못했습니다',
|
||||
rotateSecret: '비밀 키 교체',
|
||||
rotateHint: '새 서명 비밀 키를 발급합니다. 서비스에 새 키를 설정하기 전까지 호출이 실패합니다.',
|
||||
rotateConfirm: '비밀 키를 교체할까요? 이전 키는 즉시 무효화됩니다.',
|
||||
rotateFailed: '비밀 키를 교체하지 못했습니다'
|
||||
}
|
||||
},
|
||||
pluginCenter: {
|
||||
|
||||
@@ -5254,7 +5254,8 @@ export default {
|
||||
pluginAdmin: {
|
||||
runtime: {
|
||||
declarative: 'Декларативный',
|
||||
host: 'Локальный процесс'
|
||||
host: 'Локальный процесс',
|
||||
remote: 'Удалённый сервис'
|
||||
},
|
||||
title: 'Управление плагинами',
|
||||
description: 'Установка и управление плагинами всей платформы. Установленный плагин виден всем рабочим пространствам, и каждое включает его самостоятельно; отключение здесь выгружает плагин на всех узлах.',
|
||||
@@ -5298,6 +5299,13 @@ export default {
|
||||
hostApi: 'API WeKnora',
|
||||
events: 'События'
|
||||
},
|
||||
secret: {
|
||||
title: 'Секрет подписи плагина',
|
||||
description: 'Укажите этот секрет в переменной окружения WEKNORA_PLUGIN_SECRET сервиса плагина: по нему сервис проверяет, что запросы приходят от WeKnora. Секрет показывается только один раз и после закрытия недоступен.',
|
||||
copy: 'Копировать',
|
||||
copied: 'Скопировано',
|
||||
done: 'Я сохранил его'
|
||||
},
|
||||
install: {
|
||||
title: 'Установка плагина',
|
||||
description: 'Загрузите пакет .wkp или укажите URL для скачивания, проверьте и установите.',
|
||||
@@ -5312,6 +5320,9 @@ export default {
|
||||
fileHint: 'Файл .wkp (zip с plugin.yaml), до 64 МБ.',
|
||||
urlLabel: 'URL для скачивания',
|
||||
urlHint: 'Сервер скачает этот URL; адреса частных сетей запрещены.',
|
||||
remoteUrlLabel: 'URL сервиса',
|
||||
remoteUrlHint: 'HTTP(S)-адрес сервиса плагина. Хосты частной сети должны быть указаны в SSRF_WHITELIST.',
|
||||
remoteUrlKeep: 'При обновлении оставьте пустым, чтобы сохранить текущий URL',
|
||||
reviewSection: 'Проверка перед установкой',
|
||||
noPermissions: 'Плагин не запрашивает дополнительных разрешений.',
|
||||
configNotice: 'Плагину нужны настройки: настройки платформы — в его карточке после установки, настройки рабочего пространства заполняют его администраторы в разделе «Плагины».',
|
||||
@@ -5354,7 +5365,16 @@ export default {
|
||||
uninstall: 'Удалить плагин',
|
||||
uninstallConfirm: 'Удалить? Все узлы сразу выгрузят плагин.',
|
||||
uninstalled: 'Плагин удалён',
|
||||
uninstallFailed: 'Не удалось удалить'
|
||||
uninstallFailed: 'Не удалось удалить',
|
||||
remote: 'Удалённый сервис',
|
||||
remoteUrl: 'URL сервиса',
|
||||
editUrl: 'Изменить',
|
||||
urlSaved: 'URL сервиса обновлён',
|
||||
urlSaveFailed: 'Не удалось обновить URL сервиса',
|
||||
rotateSecret: 'Сменить секрет',
|
||||
rotateHint: 'Выпускает новый секрет подписи. Вызовы сервиса будут завершаться ошибкой, пока он не получит новый секрет.',
|
||||
rotateConfirm: 'Сменить секрет? Старый сразу перестанет действовать.',
|
||||
rotateFailed: 'Не удалось сменить секрет'
|
||||
}
|
||||
},
|
||||
pluginCenter: {
|
||||
|
||||
@@ -5256,7 +5256,8 @@ export default {
|
||||
pluginAdmin: {
|
||||
runtime: {
|
||||
declarative: '声明式',
|
||||
host: '本机进程'
|
||||
host: '本机进程',
|
||||
remote: '远程服务'
|
||||
},
|
||||
title: '插件管理',
|
||||
description: '安装与管理本平台的插件。安装后所有空间都能看到,但每个空间需自行启用;停用会在所有节点卸载该插件。',
|
||||
@@ -5300,6 +5301,13 @@ export default {
|
||||
hostApi: '调用 WeKnora API',
|
||||
events: '订阅事件'
|
||||
},
|
||||
secret: {
|
||||
title: '插件签名密钥',
|
||||
description: '把此密钥配置为插件服务的环境变量 WEKNORA_PLUGIN_SECRET,服务据此确认请求来自 WeKnora。密钥只显示这一次,关闭后无法再查看。',
|
||||
copy: '复制',
|
||||
copied: '已复制',
|
||||
done: '我已保存'
|
||||
},
|
||||
install: {
|
||||
title: '安装插件',
|
||||
description: '上传 .wkp 插件包或填写下载地址,审阅后安装。',
|
||||
@@ -5314,6 +5322,9 @@ export default {
|
||||
fileHint: '.wkp 文件(包含 plugin.yaml 的 zip),最大 64 MB。',
|
||||
urlLabel: '下载地址',
|
||||
urlHint: '服务端会下载该地址;不允许内网地址。',
|
||||
remoteUrlLabel: '服务地址',
|
||||
remoteUrlHint: '插件服务的 HTTP(S) 地址。内网地址需要加入 SSRF_WHITELIST。',
|
||||
remoteUrlKeep: '升级时可留空,沿用当前地址',
|
||||
reviewSection: '安装前审阅',
|
||||
noPermissions: '该插件不申请任何额外权限。',
|
||||
configNotice: '该插件需要配置:平台配置在安装后的详情中填写,空间配置由各空间管理员在插件中心填写。',
|
||||
@@ -5356,7 +5367,16 @@ export default {
|
||||
uninstall: '卸载插件',
|
||||
uninstallConfirm: '确定卸载?所有节点会立即卸载该插件。',
|
||||
uninstalled: '插件已卸载',
|
||||
uninstallFailed: '卸载失败'
|
||||
uninstallFailed: '卸载失败',
|
||||
remote: '远程服务',
|
||||
remoteUrl: '服务地址',
|
||||
editUrl: '修改',
|
||||
urlSaved: '服务地址已更新',
|
||||
urlSaveFailed: '更新服务地址失败',
|
||||
rotateSecret: '轮换密钥',
|
||||
rotateHint: '生成新的签名密钥。服务换上新密钥之前,对它的调用会失败。',
|
||||
rotateConfirm: '确定轮换?旧密钥立即失效。',
|
||||
rotateFailed: '轮换密钥失败'
|
||||
}
|
||||
},
|
||||
pluginCenter: {
|
||||
|
||||
@@ -64,6 +64,18 @@
|
||||
</div>
|
||||
|
||||
<PluginInstallDrawer v-model:visible="installOpen" @installed="upsert" />
|
||||
<t-dialog
|
||||
v-model:visible="secretOpen"
|
||||
:header="t('pluginAdmin.secret.title')"
|
||||
:confirm-btn="{ content: t('pluginAdmin.secret.copy'), theme: 'primary' }"
|
||||
:cancel-btn="{ content: t('pluginAdmin.secret.done') }"
|
||||
:close-on-overlay-click="false"
|
||||
@confirm="copyWithToast(issuedSecret, 'pluginAdmin.secret.copied')"
|
||||
@closed="issuedSecret = ''"
|
||||
>
|
||||
<p>{{ t('pluginAdmin.secret.description') }}</p>
|
||||
<t-textarea :value="issuedSecret" readonly autosize />
|
||||
</t-dialog>
|
||||
<PluginDetailDrawer
|
||||
v-model:visible="detailOpen"
|
||||
:plugin="selected"
|
||||
@@ -79,6 +91,7 @@ import { useI18n } from 'vue-i18n'
|
||||
import { MessagePlugin } from 'tdesign-vue-next'
|
||||
|
||||
import { listInstalledPlugins, setInstalledPluginEnabled, type InstalledPlugin } from '@/api/system/plugins'
|
||||
import { copyWithToast } from '@/utils/clipboard'
|
||||
import { localizedText } from '@/utils/localizedText'
|
||||
|
||||
import { contributionSummary } from '../settings/pluginCenterState'
|
||||
@@ -97,6 +110,9 @@ const loading = ref(false)
|
||||
const pending = ref(new Set<string>())
|
||||
const installOpen = ref(false)
|
||||
const detailOpen = ref(false)
|
||||
// A remote plugin's signing secret, shown once after it is issued.
|
||||
const secretOpen = ref(false)
|
||||
const issuedSecret = ref('')
|
||||
const selectedId = ref('')
|
||||
const selected = computed(() => plugins.value.find((p) => p.id === selectedId.value) ?? null)
|
||||
|
||||
@@ -129,7 +145,12 @@ async function load() {
|
||||
}
|
||||
}
|
||||
|
||||
function upsert(p: InstalledPlugin) {
|
||||
function upsert(received: InstalledPlugin) {
|
||||
const { issuedSecret: secret, ...p } = received
|
||||
if (secret) {
|
||||
issuedSecret.value = secret
|
||||
secretOpen.value = true
|
||||
}
|
||||
const i = plugins.value.findIndex((x) => x.id === p.id)
|
||||
if (i >= 0) plugins.value.splice(i, 1, p)
|
||||
else plugins.value = [...plugins.value, p].sort((a, b) => a.id.localeCompare(b.id))
|
||||
|
||||
@@ -12,6 +12,7 @@ import {
|
||||
isPackageUrl,
|
||||
permissionLines,
|
||||
remoteHosts,
|
||||
remoteUrlReady,
|
||||
shortDigest,
|
||||
sortVersions,
|
||||
} from './pluginManagementState'
|
||||
@@ -77,3 +78,14 @@ test('isPackageUrl and hasSystemConfig', () => {
|
||||
assert.equal(hasSystemConfig(manifest), true)
|
||||
assert.equal(hasSystemConfig(undefined), false)
|
||||
})
|
||||
|
||||
test('remoteUrlReady asks a new remote plugin for its service URL', () => {
|
||||
const remote = (change: string) => ({ manifest: { runtime: { type: 'remote' } }, change })
|
||||
assert.equal(remoteUrlReady(remote('install'), ''), false)
|
||||
assert.equal(remoteUrlReady(remote('install'), 'plugins.example.com'), false)
|
||||
assert.equal(remoteUrlReady(remote('install'), 'https://plugins.example.com'), true)
|
||||
assert.equal(remoteUrlReady(remote('upgrade'), ''), true)
|
||||
assert.equal(remoteUrlReady(remote('upgrade'), 'nope'), false)
|
||||
assert.equal(remoteUrlReady({ manifest: { runtime: { type: 'host' } }, change: 'install' }, ''), true)
|
||||
assert.equal(remoteUrlReady(null, ''), true)
|
||||
})
|
||||
|
||||
@@ -114,6 +114,16 @@ export function installedState(p: InstalledPlugin): InstalledState {
|
||||
export const EGRESS_ANY_HOST = '*'
|
||||
|
||||
/** Whether a string can be sent as a package URL. */
|
||||
/**
|
||||
* Whether the install drawer can go ahead with a remote plugin's service URL:
|
||||
* a new install needs one, an upgrade may keep the registered one.
|
||||
*/
|
||||
export function remoteUrlReady(preview: { manifest: { runtime?: { type: string } }; change: string } | null, raw: string) {
|
||||
if (preview?.manifest.runtime?.type !== 'remote') return true
|
||||
if (raw.trim() === '') return preview.change !== 'install'
|
||||
return isPackageUrl(raw)
|
||||
}
|
||||
|
||||
export function isPackageUrl(raw: string): boolean {
|
||||
try {
|
||||
const u = new URL(raw.trim())
|
||||
|
||||
@@ -63,6 +63,34 @@
|
||||
</ul>
|
||||
</section>
|
||||
|
||||
<section v-if="plugin.runtime === 'remote'" class="setting-drawer__section">
|
||||
<h4 class="setting-drawer__section-title">{{ t('pluginAdmin.detail.remote') }}</h4>
|
||||
<div class="remote-row">
|
||||
<span class="remote-row__label">{{ t('pluginAdmin.detail.remoteUrl') }}</span>
|
||||
<template v-if="editingUrl">
|
||||
<t-input v-model="urlDraft" size="small" class="remote-row__input" :disabled="savingUrl" @enter="saveUrl" />
|
||||
<t-button size="small" theme="primary" :loading="savingUrl" :disabled="!isPackageUrl(urlDraft)" @click="saveUrl">
|
||||
{{ t('common.save') }}
|
||||
</t-button>
|
||||
<t-button size="small" variant="text" :disabled="savingUrl" @click="editingUrl = false">
|
||||
{{ t('common.cancel') }}
|
||||
</t-button>
|
||||
</template>
|
||||
<template v-else>
|
||||
<code class="remote-row__value">{{ plugin.remote_url }}</code>
|
||||
<t-button size="small" variant="text" theme="primary" @click="startEditUrl">
|
||||
{{ t('pluginAdmin.detail.editUrl') }}
|
||||
</t-button>
|
||||
</template>
|
||||
</div>
|
||||
<div class="danger-row">
|
||||
<span class="form-desc">{{ t('pluginAdmin.detail.rotateHint') }}</span>
|
||||
<t-popconfirm :content="t('pluginAdmin.detail.rotateConfirm')" @confirm="rotate">
|
||||
<t-button variant="outline" :loading="rotating">{{ t('pluginAdmin.detail.rotateSecret') }}</t-button>
|
||||
</t-popconfirm>
|
||||
</div>
|
||||
</section>
|
||||
|
||||
<section class="setting-drawer__section">
|
||||
<h4 class="setting-drawer__section-title">{{ t('pluginAdmin.detail.versions') }}</h4>
|
||||
<ul class="line-list">
|
||||
@@ -119,6 +147,8 @@ import { getPlugin, type PluginInstance } from '@/api/plugin'
|
||||
import {
|
||||
activatePluginVersion,
|
||||
getPluginSystemConfig,
|
||||
rotatePluginSecret,
|
||||
setPluginRemoteUrl,
|
||||
uninstallPlugin,
|
||||
updatePluginSystemConfig,
|
||||
type InstalledPlugin,
|
||||
@@ -130,6 +160,7 @@ import {
|
||||
contributionLines,
|
||||
formatBytes,
|
||||
hasSystemConfig,
|
||||
isPackageUrl,
|
||||
shortDigest,
|
||||
sortVersions,
|
||||
} from '../pluginManagementState'
|
||||
@@ -156,6 +187,10 @@ const configErrors = ref<FieldError[]>([])
|
||||
const saving = ref(false)
|
||||
const activating = ref('')
|
||||
const uninstalling = ref(false)
|
||||
const editingUrl = ref(false)
|
||||
const urlDraft = ref('')
|
||||
const savingUrl = ref(false)
|
||||
const rotating = ref(false)
|
||||
|
||||
const formatDate = (s: string) => (s ? new Date(s).toLocaleString(locale.value) : '')
|
||||
|
||||
@@ -188,6 +223,7 @@ watch(
|
||||
() => [props.visible, props.plugin?.id, props.plugin?.active_version, props.plugin?.desired_state] as const,
|
||||
([visible, id]) => {
|
||||
if (!visible || !id) return
|
||||
editingUrl.value = false
|
||||
void loadNodes(id)
|
||||
void loadConfig(id)
|
||||
},
|
||||
@@ -226,6 +262,40 @@ async function activate(version: string) {
|
||||
}
|
||||
}
|
||||
|
||||
function startEditUrl() {
|
||||
urlDraft.value = props.plugin?.remote_url ?? ''
|
||||
editingUrl.value = true
|
||||
}
|
||||
|
||||
async function saveUrl() {
|
||||
if (!props.plugin || !isPackageUrl(urlDraft.value)) return
|
||||
savingUrl.value = true
|
||||
try {
|
||||
const res = await setPluginRemoteUrl(props.plugin.id, urlDraft.value.trim())
|
||||
editingUrl.value = false
|
||||
emit('changed', res.data)
|
||||
MessagePlugin.success(t('pluginAdmin.detail.urlSaved'))
|
||||
} catch (e: any) {
|
||||
MessagePlugin.error(e?.message || t('pluginAdmin.detail.urlSaveFailed'))
|
||||
} finally {
|
||||
savingUrl.value = false
|
||||
}
|
||||
}
|
||||
|
||||
// The new secret travels in the changed plugin; the page shows it once.
|
||||
async function rotate() {
|
||||
if (!props.plugin) return
|
||||
rotating.value = true
|
||||
try {
|
||||
const res = await rotatePluginSecret(props.plugin.id)
|
||||
emit('changed', res.data)
|
||||
} catch (e: any) {
|
||||
MessagePlugin.error(e?.message || t('pluginAdmin.detail.rotateFailed'))
|
||||
} finally {
|
||||
rotating.value = false
|
||||
}
|
||||
}
|
||||
|
||||
async function uninstall() {
|
||||
if (!props.plugin) return
|
||||
const id = props.plugin.id
|
||||
@@ -318,6 +388,29 @@ async function uninstall() {
|
||||
color: var(--td-text-color-placeholder);
|
||||
}
|
||||
|
||||
.remote-row {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
flex-wrap: wrap;
|
||||
gap: 8px;
|
||||
font-size: var(--app-text-sm);
|
||||
|
||||
&__label {
|
||||
color: var(--td-text-color-secondary);
|
||||
}
|
||||
|
||||
&__value {
|
||||
font-size: var(--app-text-xs);
|
||||
color: var(--td-text-color-primary);
|
||||
word-break: break-all;
|
||||
}
|
||||
|
||||
&__input {
|
||||
flex: 1;
|
||||
min-width: 200px;
|
||||
}
|
||||
}
|
||||
|
||||
.danger-row {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
|
||||
@@ -104,6 +104,17 @@
|
||||
</ul>
|
||||
</div>
|
||||
|
||||
<div v-if="isRemote" class="form-item">
|
||||
<label class="form-label" :class="{ required: preview.change === 'install' }">
|
||||
{{ t('pluginAdmin.install.remoteUrlLabel') }}
|
||||
</label>
|
||||
<t-input v-model="remoteUrl" :disabled="busy" placeholder="https://plugins.example.com/acme-search" />
|
||||
<p class="form-desc">
|
||||
{{ t('pluginAdmin.install.remoteUrlHint') }}
|
||||
<template v-if="preview.change !== 'install'">{{ t('pluginAdmin.install.remoteUrlKeep') }}</template>
|
||||
</p>
|
||||
</div>
|
||||
|
||||
<t-alert
|
||||
v-if="needsConfig"
|
||||
theme="info"
|
||||
@@ -139,6 +150,7 @@ import {
|
||||
isPackageUrl,
|
||||
permissionLines,
|
||||
remoteHosts,
|
||||
remoteUrlReady,
|
||||
} from '../pluginManagementState'
|
||||
import { hasTenantConfig } from '../../settings/pluginCenterState'
|
||||
|
||||
@@ -156,6 +168,7 @@ const { t, locale } = useI18n()
|
||||
const mode = ref<'upload' | 'url'>('upload')
|
||||
const file = ref<File | null>(null)
|
||||
const url = ref('')
|
||||
const remoteUrl = ref('')
|
||||
const fileInput = ref<HTMLInputElement | null>(null)
|
||||
const preview = ref<PluginPreview | null>(null)
|
||||
const busy = ref(false)
|
||||
@@ -167,6 +180,7 @@ watch(
|
||||
mode.value = 'upload'
|
||||
file.value = null
|
||||
url.value = ''
|
||||
remoteUrl.value = ''
|
||||
preview.value = null
|
||||
},
|
||||
)
|
||||
@@ -189,7 +203,10 @@ const source = computed<PackageSource | null>(() => {
|
||||
return isPackageUrl(url.value) ? { url: url.value.trim() } : null
|
||||
})
|
||||
|
||||
const canConfirm = computed(() => !!source.value && !busy.value)
|
||||
const isRemote = computed(() => preview.value?.manifest.runtime?.type === 'remote')
|
||||
const canConfirm = computed(
|
||||
() => !!source.value && !busy.value && (!preview.value || remoteUrlReady(preview.value, remoteUrl.value)),
|
||||
)
|
||||
const confirmText = computed(() =>
|
||||
preview.value ? t(`pluginAdmin.install.confirm.${preview.value.change}`) : t('pluginAdmin.install.inspect'),
|
||||
)
|
||||
@@ -237,7 +254,11 @@ async function onConfirm() {
|
||||
if (!source.value) return
|
||||
busy.value = true
|
||||
try {
|
||||
const res = await installPluginPackage(source.value, preview.value.digest)
|
||||
const res = await installPluginPackage(
|
||||
source.value,
|
||||
preview.value.digest,
|
||||
isRemote.value ? remoteUrl.value.trim() : undefined,
|
||||
)
|
||||
MessagePlugin.success(t('pluginAdmin.install.done', { name: localizedText(preview.value.manifest.name, locale.value) }))
|
||||
emit('installed', res.data)
|
||||
emit('update:visible', false)
|
||||
|
||||
@@ -40,13 +40,17 @@ func (r *pluginRepository) GetPlugin(ctx context.Context, id string) (*types.Ins
|
||||
}
|
||||
|
||||
// SavePlugin inserts or updates the row keyed by id.
|
||||
// pluginUpdateColumns are what saving an installed plugin overwrites: every
|
||||
// column but its identity and creation.
|
||||
var pluginUpdateColumns = []string{
|
||||
"owner_tenant_id", "source", "active_version", "desired_state", "runtime",
|
||||
"granted_perms", "system_config", "remote_url", "remote_secret", "updated_at",
|
||||
}
|
||||
|
||||
func (r *pluginRepository) SavePlugin(ctx context.Context, p *types.InstalledPlugin) error {
|
||||
return r.db.WithContext(ctx).Clauses(clause.OnConflict{
|
||||
Columns: []clause.Column{{Name: "id"}},
|
||||
DoUpdates: clause.AssignmentColumns([]string{
|
||||
"owner_tenant_id", "source", "active_version", "desired_state", "runtime",
|
||||
"granted_perms", "system_config", "updated_at",
|
||||
}),
|
||||
Columns: []clause.Column{{Name: "id"}},
|
||||
DoUpdates: clause.AssignmentColumns(pluginUpdateColumns),
|
||||
}).Create(p).Error
|
||||
}
|
||||
|
||||
|
||||
@@ -2,6 +2,7 @@ package repository
|
||||
|
||||
import (
|
||||
"context"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
@@ -9,6 +10,7 @@ import (
|
||||
"github.com/stretchr/testify/require"
|
||||
"gorm.io/driver/sqlite"
|
||||
"gorm.io/gorm"
|
||||
"gorm.io/gorm/schema"
|
||||
|
||||
"github.com/Tencent/WeKnora/internal/types"
|
||||
)
|
||||
@@ -39,12 +41,15 @@ func TestPluginRepository(t *testing.T) {
|
||||
require.NoError(t, repo.SavePlugin(ctx, p))
|
||||
p.ActiveVersion = "1.1.0"
|
||||
p.DesiredState = types.PluginStateDisabled
|
||||
p.RemoteURL, p.RemoteSecret = "https://plugins.example.com", "enc:v1:x"
|
||||
require.NoError(t, repo.SavePlugin(ctx, p), "save is an upsert")
|
||||
|
||||
got, err = repo.GetPlugin(ctx, "acme.kit")
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "1.1.0", got.ActiveVersion)
|
||||
require.Equal(t, types.PluginStateDisabled, got.DesiredState)
|
||||
require.Equal(t, "https://plugins.example.com", got.RemoteURL)
|
||||
require.Equal(t, "enc:v1:x", got.RemoteSecret)
|
||||
|
||||
versions, err := repo.ListVersions(ctx, "acme.kit")
|
||||
require.NoError(t, err)
|
||||
@@ -142,3 +147,17 @@ func TestPluginKVRepository(t *testing.T) {
|
||||
got, _ = repo.Get(ctx, "acme.x", 2, "cursor:a")
|
||||
require.JSONEq(t, `99`, string(got.Value), "tenants are separate partitions")
|
||||
}
|
||||
|
||||
// A column added to installed plugins must be saved on update too, or
|
||||
// changing it silently does nothing.
|
||||
func TestPluginUpdateColumnsCoverTheTable(t *testing.T) {
|
||||
s, err := schema.Parse(&types.InstalledPlugin{}, &sync.Map{}, schema.NamingStrategy{})
|
||||
require.NoError(t, err)
|
||||
for _, name := range s.DBNames {
|
||||
switch name {
|
||||
case "id", "created_at", "created_by":
|
||||
continue
|
||||
}
|
||||
require.Contains(t, pluginUpdateColumns, name, "SavePlugin does not update %s", name)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -78,6 +78,7 @@ import (
|
||||
pluginmanifest "github.com/Tencent/WeKnora/internal/plugin/manifest"
|
||||
"github.com/Tencent/WeKnora/internal/plugin/reconcile"
|
||||
pluginregistry "github.com/Tencent/WeKnora/internal/plugin/registry"
|
||||
pluginremote "github.com/Tencent/WeKnora/internal/plugin/remote"
|
||||
plugintenancy "github.com/Tencent/WeKnora/internal/plugin/tenancy"
|
||||
"github.com/Tencent/WeKnora/internal/router"
|
||||
"github.com/Tencent/WeKnora/internal/sandbox"
|
||||
@@ -168,6 +169,7 @@ func BuildContainer(container *dig.Container) *dig.Container {
|
||||
must(container.Provide(activate.NewSkills))
|
||||
must(container.Provide(activate.NewModelVendors))
|
||||
must(container.Provide(pluginhost.NewManager))
|
||||
must(container.Provide(pluginremote.NewManager))
|
||||
must(container.Provide(newPluginInvoker))
|
||||
must(container.Provide(repository.NewPluginKVRepository))
|
||||
must(container.Provide(newPluginHostAPI))
|
||||
@@ -1840,12 +1842,13 @@ func newPluginRegistry(
|
||||
})
|
||||
}
|
||||
|
||||
// newPluginDrivers returns the drivers that run plugins: builtins and
|
||||
// declarative packages today; host, remote and kubernetes drivers join this
|
||||
// set.
|
||||
// newPluginDrivers returns the drivers that run plugins: builtins,
|
||||
// declarative packages, the embedded host and remote services; a kubernetes
|
||||
// driver joins this set.
|
||||
func newPluginDrivers(reg *pluginregistry.Registry, r *reconcile.Reconciler) *plugindriver.Set {
|
||||
return plugindriver.NewSet(plugindriver.NewBuiltin(reg.Plugin),
|
||||
r.Driver(pluginmanifest.RuntimeDeclarative), r.Driver(pluginmanifest.RuntimeHost))
|
||||
r.Driver(pluginmanifest.RuntimeDeclarative), r.Driver(pluginmanifest.RuntimeHost),
|
||||
r.Driver(pluginmanifest.RuntimeRemote))
|
||||
}
|
||||
|
||||
// installPluginGate gives the integration handlers the tenant plugin switches,
|
||||
|
||||
@@ -21,8 +21,10 @@ import (
|
||||
"github.com/Tencent/WeKnora/internal/plugin/install"
|
||||
"github.com/Tencent/WeKnora/internal/plugin/reconcile"
|
||||
pluginregistry "github.com/Tencent/WeKnora/internal/plugin/registry"
|
||||
"github.com/Tencent/WeKnora/internal/plugin/remote"
|
||||
plugintenancy "github.com/Tencent/WeKnora/internal/plugin/tenancy"
|
||||
"github.com/Tencent/WeKnora/internal/types/interfaces"
|
||||
"github.com/Tencent/WeKnora/pluginsdk/client"
|
||||
)
|
||||
|
||||
// newPluginPackageStore keeps plugin packages in the deployment's object
|
||||
@@ -41,6 +43,7 @@ type pluginActivators struct {
|
||||
dig.In
|
||||
|
||||
Host *host.Manager
|
||||
Remote *remote.Manager
|
||||
WebSearch *activate.WebSearch
|
||||
Connectors *activate.Connectors
|
||||
Parsers *activate.Parsers
|
||||
@@ -50,19 +53,23 @@ type pluginActivators struct {
|
||||
Invoker *activate.Invoker
|
||||
}
|
||||
|
||||
// list orders the activators: the host first, so a code plugin's process is
|
||||
// running before anything routes calls to it.
|
||||
// list orders the activators: the runtimes first, so a code plugin is
|
||||
// reachable before anything routes calls to it.
|
||||
func (a pluginActivators) list() []reconcile.Activator {
|
||||
return []reconcile.Activator{a.Host, a.WebSearch, a.Connectors, a.Parsers, a.Vendors, a.MCP, a.Skills}
|
||||
return []reconcile.Activator{
|
||||
a.Host, a.Remote, a.WebSearch, a.Connectors, a.Parsers, a.Vendors, a.MCP, a.Skills,
|
||||
}
|
||||
}
|
||||
|
||||
// newPluginHostAPI serves the Host API and gives calls a way back to it: the
|
||||
// embedded host's plugins reach this node on loopback.
|
||||
// embedded host's plugins reach this node on loopback, remote plugins only
|
||||
// through WEKNORA_PLUGIN_HOST_API_URL.
|
||||
func newPluginHostAPI(
|
||||
cfg *config.Config, iv *activate.Invoker, repo interfaces.PluginKVRepository, cleaner interfaces.ResourceCleaner,
|
||||
) *hostapi.Handler {
|
||||
issuer := hostapi.NewIssuerFromEnv()
|
||||
url := strings.TrimSpace(os.Getenv("WEKNORA_PLUGIN_HOST_API_URL"))
|
||||
iv.SetPublicHostAPI(url)
|
||||
if url == "" {
|
||||
port := 8080
|
||||
if cfg != nil && cfg.Server != nil && cfg.Server.Port > 0 {
|
||||
@@ -90,8 +97,24 @@ func bindPluginActivators(
|
||||
skills.SetPluginSkills(a.Skills)
|
||||
}
|
||||
|
||||
// newPluginInvoker reaches code plugins through this node's plugin host.
|
||||
func newPluginInvoker(h *host.Manager) *activate.Invoker { return activate.NewInvoker(h) }
|
||||
// newPluginInvoker reaches code plugins through this node's plugin host or,
|
||||
// for remote ones, at their registered URLs.
|
||||
func newPluginInvoker(h *host.Manager, r *remote.Manager) *activate.Invoker {
|
||||
return activate.NewInvoker(pluginClients{host: h, remote: r})
|
||||
}
|
||||
|
||||
// pluginClients finds a code plugin in whichever runtime loaded it.
|
||||
type pluginClients struct {
|
||||
host *host.Manager
|
||||
remote *remote.Manager
|
||||
}
|
||||
|
||||
func (c pluginClients) Client(pluginID string) (*client.Client, error) {
|
||||
if c.remote.Owns(pluginID) {
|
||||
return c.remote.Client(pluginID)
|
||||
}
|
||||
return c.host.Client(pluginID)
|
||||
}
|
||||
|
||||
func newMCPServiceRepository(db *gorm.DB, plugins *activate.MCPServers) interfaces.MCPServiceRepository {
|
||||
return plugins.Repository(repository.NewMCPServiceRepository(db))
|
||||
@@ -119,13 +142,18 @@ func newPluginInstaller(
|
||||
|
||||
// startPluginReconciler loads installed plugins before the server takes
|
||||
// traffic, then keeps this node in step with the others.
|
||||
func startPluginReconciler(r *reconcile.Reconciler, hostManager *host.Manager, cleaner interfaces.ResourceCleaner) {
|
||||
func startPluginReconciler(
|
||||
r *reconcile.Reconciler, hostManager *host.Manager, remoteManager *remote.Manager,
|
||||
cleaner interfaces.ResourceCleaner,
|
||||
) {
|
||||
hostManager.SetReporter(r)
|
||||
remoteManager.SetReporter(r)
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
r.Start(ctx)
|
||||
cleaner.RegisterWithName("PluginReconciler", func() error {
|
||||
cancel()
|
||||
hostManager.Close()
|
||||
remoteManager.Close()
|
||||
return nil
|
||||
})
|
||||
}
|
||||
|
||||
@@ -73,7 +73,7 @@ var versionedSQLiteColumns = map[string][]string{
|
||||
}, // 000028
|
||||
}
|
||||
|
||||
const expectedSQLiteMigrationVersion = 36
|
||||
const expectedSQLiteMigrationVersion = 37
|
||||
|
||||
func TestSQLiteMigrationsCreateVersionedSchema(t *testing.T) {
|
||||
repoRoot := sqliteRepoRoot(t)
|
||||
|
||||
@@ -32,23 +32,36 @@ type PluginPackageRequest struct {
|
||||
URL string `json:"url"`
|
||||
// Digest pins an install to the package reviewed with inspect.
|
||||
Digest string `json:"digest"`
|
||||
// RemoteURL is where a remote plugin's service runs.
|
||||
RemoteURL string `json:"remote_url"`
|
||||
}
|
||||
|
||||
// packageInput is a package with what its install request said about it.
|
||||
type packageInput struct {
|
||||
data []byte
|
||||
source install.Source
|
||||
digest string
|
||||
remoteURL string
|
||||
}
|
||||
|
||||
// readPackage returns the package archive from an upload or a URL.
|
||||
func (h *PluginAdminHandler) readPackage(c *gin.Context) ([]byte, install.Source, string, bool) {
|
||||
func (h *PluginAdminHandler) readPackage(c *gin.Context) (packageInput, bool) {
|
||||
if strings.HasPrefix(c.ContentType(), "application/json") {
|
||||
limitJSONBody(c, skillSourceJSONMaxBytes)
|
||||
var req PluginPackageRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil || strings.TrimSpace(req.URL) == "" {
|
||||
_ = c.Error(errors.NewBadRequestError("url or an uploaded file is required"))
|
||||
return nil, install.Source{}, "", false
|
||||
return packageInput{}, false
|
||||
}
|
||||
data, err := h.service.FetchURL(c.Request.Context(), strings.TrimSpace(req.URL))
|
||||
if err != nil {
|
||||
h.fail(c, err)
|
||||
return nil, install.Source{}, "", false
|
||||
return packageInput{}, false
|
||||
}
|
||||
return data, install.Source{Kind: "url", URL: strings.TrimSpace(req.URL)}, req.Digest, true
|
||||
return packageInput{
|
||||
data: data, source: install.Source{Kind: "url", URL: strings.TrimSpace(req.URL)},
|
||||
digest: req.Digest, remoteURL: strings.TrimSpace(req.RemoteURL),
|
||||
}, true
|
||||
}
|
||||
|
||||
limitUploadBody(c, pkg.MaxArchiveBytes)
|
||||
@@ -59,19 +72,22 @@ func (h *PluginAdminHandler) readPackage(c *gin.Context) ([]byte, install.Source
|
||||
} else {
|
||||
_ = c.Error(errors.NewBadRequestError("file is required"))
|
||||
}
|
||||
return nil, install.Source{}, "", false
|
||||
return packageInput{}, false
|
||||
}
|
||||
defer func() { _ = file.Close() }()
|
||||
data, err := io.ReadAll(io.LimitReader(file, pkg.MaxArchiveBytes+1))
|
||||
if err != nil {
|
||||
_ = c.Error(errors.NewBadRequestError("failed to read the uploaded package"))
|
||||
return nil, install.Source{}, "", false
|
||||
return packageInput{}, false
|
||||
}
|
||||
if len(data) > pkg.MaxArchiveBytes {
|
||||
_ = c.Error(errors.NewBadRequestError("plugin package is too large"))
|
||||
return nil, install.Source{}, "", false
|
||||
return packageInput{}, false
|
||||
}
|
||||
return data, install.Source{Kind: "upload", URL: header.Filename}, c.PostForm("digest"), true
|
||||
return packageInput{
|
||||
data: data, source: install.Source{Kind: "upload", URL: header.Filename},
|
||||
digest: c.PostForm("digest"), remoteURL: strings.TrimSpace(c.PostForm("remote_url")),
|
||||
}, true
|
||||
}
|
||||
|
||||
func (h *PluginAdminHandler) fail(c *gin.Context, err error) {
|
||||
@@ -101,11 +117,11 @@ func (h *PluginAdminHandler) ok(c *gin.Context, data any) {
|
||||
// @Security Bearer
|
||||
// @Router /system/admin/plugins/inspect [post]
|
||||
func (h *PluginAdminHandler) InspectPlugin(c *gin.Context) {
|
||||
data, _, _, ok := h.readPackage(c)
|
||||
in, ok := h.readPackage(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
preview, err := h.service.Inspect(c.Request.Context(), data)
|
||||
preview, err := h.service.Inspect(c.Request.Context(), in.data)
|
||||
if err != nil {
|
||||
h.fail(c, err)
|
||||
return
|
||||
@@ -115,7 +131,9 @@ func (h *PluginAdminHandler) InspectPlugin(c *gin.Context) {
|
||||
|
||||
// InstallPlugin godoc
|
||||
// @Summary 安装或升级插件
|
||||
// @Description 安装插件包(multipart file + digest,或 JSON {url, digest})。digest 取自 inspect,保证安装的就是审阅过的包。安装后全平台可用,各空间需自行启用
|
||||
// @Description 安装插件包(multipart file + digest,或 JSON {url, digest})。digest 取自 inspect,
|
||||
// @Description 保证安装的就是审阅过的包。安装后全平台可用,各空间需自行启用。
|
||||
// @Description remote 插件还需 remote_url(服务地址);首次安装的响应里 issuedSecret 是签名密钥,只返回这一次
|
||||
// @Tags System
|
||||
// @Accept multipart/form-data,json
|
||||
// @Produce json
|
||||
@@ -123,13 +141,13 @@ func (h *PluginAdminHandler) InspectPlugin(c *gin.Context) {
|
||||
// @Security Bearer
|
||||
// @Router /system/admin/plugins [post]
|
||||
func (h *PluginAdminHandler) InstallPlugin(c *gin.Context) {
|
||||
data, source, digest, ok := h.readPackage(c)
|
||||
in, ok := h.readPackage(c)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
userID, _ := c.Request.Context().Value(types.UserIDContextKey).(string)
|
||||
view, err := h.service.Install(c.Request.Context(), install.Request{
|
||||
Data: data, Source: source, ExpectedDigest: digest, UserID: userID,
|
||||
Data: in.data, Source: in.source, ExpectedDigest: in.digest, UserID: userID, RemoteURL: in.remoteURL,
|
||||
})
|
||||
if err != nil {
|
||||
h.fail(c, err)
|
||||
@@ -227,6 +245,54 @@ func (h *PluginAdminHandler) ActivatePluginVersion(c *gin.Context) {
|
||||
h.ok(c, view)
|
||||
}
|
||||
|
||||
// SetPluginRemoteURLRequest moves a remote plugin.
|
||||
type SetPluginRemoteURLRequest struct {
|
||||
URL string `json:"url" binding:"required"`
|
||||
}
|
||||
|
||||
// SetPluginRemoteURL godoc
|
||||
// @Summary 修改远程插件的服务地址
|
||||
// @Description 私有网络中的地址需要加入 SSRF_WHITELIST
|
||||
// @Tags System
|
||||
// @Accept json
|
||||
// @Produce json
|
||||
// @Param id path string true "插件 ID"
|
||||
// @Param request body SetPluginRemoteURLRequest true "服务地址"
|
||||
// @Success 200 {object} map[string]interface{}
|
||||
// @Security Bearer
|
||||
// @Router /system/admin/plugins/{id}/remote-url [put]
|
||||
func (h *PluginAdminHandler) SetPluginRemoteURL(c *gin.Context) {
|
||||
var req SetPluginRemoteURLRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
_ = c.Error(errors.NewBadRequestError("url is required"))
|
||||
return
|
||||
}
|
||||
view, err := h.service.SetRemoteURL(c.Request.Context(), c.Param("id"), strings.TrimSpace(req.URL))
|
||||
if err != nil {
|
||||
h.fail(c, err)
|
||||
return
|
||||
}
|
||||
h.ok(c, view)
|
||||
}
|
||||
|
||||
// RotatePluginSecret godoc
|
||||
// @Summary 轮换远程插件的签名密钥
|
||||
// @Description 新密钥只在本次响应的 issuedSecret 中返回;服务换上新密钥前调用会失败
|
||||
// @Tags System
|
||||
// @Produce json
|
||||
// @Param id path string true "插件 ID"
|
||||
// @Success 200 {object} map[string]interface{}
|
||||
// @Security Bearer
|
||||
// @Router /system/admin/plugins/{id}/secret/rotate [post]
|
||||
func (h *PluginAdminHandler) RotatePluginSecret(c *gin.Context) {
|
||||
view, err := h.service.RotateSecret(c.Request.Context(), c.Param("id"))
|
||||
if err != nil {
|
||||
h.fail(c, err)
|
||||
return
|
||||
}
|
||||
h.ok(c, view)
|
||||
}
|
||||
|
||||
// UninstallPlugin godoc
|
||||
// @Summary 卸载插件
|
||||
// @Description 删除插件及其全部版本;各空间的开关和配置保留,重新安装后恢复
|
||||
|
||||
@@ -28,9 +28,11 @@ type Invoker struct {
|
||||
mu sync.RWMutex
|
||||
tenancy *tenancy.Service
|
||||
plugins interfaces.PluginRepository
|
||||
// tokens and hostURL give plugins granted Host API scopes a way back.
|
||||
tokens TokenIssuer
|
||||
hostURL string
|
||||
// tokens and hostURL give plugins granted Host API scopes a way back;
|
||||
// remote plugins, off this node, use publicHostURL.
|
||||
tokens TokenIssuer
|
||||
hostURL string
|
||||
publicHostURL string
|
||||
}
|
||||
|
||||
// TokenIssuer signs the Host API token of one call.
|
||||
@@ -46,6 +48,14 @@ func (iv *Invoker) SetHostAPI(tokens TokenIssuer, url string) {
|
||||
iv.mu.Unlock()
|
||||
}
|
||||
|
||||
// SetPublicHostAPI is where remote plugins reach the Host API. Without it
|
||||
// their calls carry no Host API access.
|
||||
func (iv *Invoker) SetPublicHostAPI(url string) {
|
||||
iv.mu.Lock()
|
||||
iv.publicHostURL = url
|
||||
iv.mu.Unlock()
|
||||
}
|
||||
|
||||
// NewInvoker creates an Invoker; Bind completes it.
|
||||
func NewInvoker(clients ClientSource) *Invoker { return &Invoker{clients: clients} }
|
||||
|
||||
@@ -94,6 +104,9 @@ func (iv *Invoker) Envelope(
|
||||
}
|
||||
iv.mu.RLock()
|
||||
t, plugins, tokens, hostURL := iv.tenancy, iv.plugins, iv.tokens, iv.hostURL
|
||||
if m.Runtime.Type == manifest.RuntimeRemote {
|
||||
hostURL = iv.publicHostURL
|
||||
}
|
||||
iv.mu.RUnlock()
|
||||
if tokens != nil && hostURL != "" && env.Context.TenantID != 0 && len(m.Permissions.HostAPI) > 0 {
|
||||
token, _, err := tokens.Issue(m.ID, m.Version, env.Context.TenantID, m.Permissions.HostAPI)
|
||||
|
||||
@@ -34,4 +34,15 @@ func TestEnvelopeCarriesHostAccessOnlyForGrantedPlugins(t *testing.T) {
|
||||
if env, _ := iv.Envelope(context.Background(), granted, nil); env.Context.Host != nil {
|
||||
t.Fatal("a call without a tenant must get no token")
|
||||
}
|
||||
|
||||
remote := *granted
|
||||
remote.Runtime.Type = manifest.RuntimeRemote
|
||||
if env, _ := iv.Envelope(ctx, &remote, nil); env.Context.Host != nil {
|
||||
t.Fatal("a remote plugin must not be sent to this node's loopback address")
|
||||
}
|
||||
iv.SetPublicHostAPI("https://weknora.example.com")
|
||||
if env, _ := iv.Envelope(ctx, &remote, nil); env.Context.Host == nil ||
|
||||
env.Context.Host.URL != "https://weknora.example.com" {
|
||||
t.Fatalf("remote envelope = %+v", env.Context)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -59,8 +59,9 @@ func (m *Manager) Activate(ctx context.Context, l *reconcile.Loaded) error {
|
||||
if l.Manifest.Runtime.Type != manifest.RuntimeHost {
|
||||
return nil
|
||||
}
|
||||
if l.Manifest.Runtime.Kind != "binary" {
|
||||
return fmt.Errorf("runtime.kind %q is not supported by this host yet; only binary", l.Manifest.Runtime.Kind)
|
||||
if !Supported(l.Manifest.Runtime.Kind) {
|
||||
return fmt.Errorf("runtime.kind %q is not supported by this host; binary and python are",
|
||||
l.Manifest.Runtime.Kind)
|
||||
}
|
||||
id := l.Manifest.ID
|
||||
p, err := startProcess(spec{m: l.Manifest, dir: l.Dir}, func(s State, err error) {
|
||||
|
||||
@@ -8,8 +8,10 @@ import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
"os"
|
||||
"os/exec"
|
||||
"path"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"strings"
|
||||
@@ -51,12 +53,40 @@ type spec struct {
|
||||
dir string // extracted package
|
||||
}
|
||||
|
||||
// entryPath resolves runtime.entry for this machine inside the package.
|
||||
func entryPath(m *manifest.Manifest, dir string) (string, error) {
|
||||
rel := strings.NewReplacer("{os}", runtime.GOOS, "{arch}", runtime.GOARCH).Replace(m.Runtime.Entry)
|
||||
if runtime.GOOS == "windows" && filepath.Ext(rel) == "" {
|
||||
// Kinds of host plugin this host runs.
|
||||
const (
|
||||
KindBinary = "binary"
|
||||
KindPython = "python"
|
||||
)
|
||||
|
||||
// Supported reports whether this host runs a kind of host plugin.
|
||||
func Supported(kind string) bool { return kind == KindBinary || kind == KindPython }
|
||||
|
||||
// PythonCommand is the interpreter python plugins run with:
|
||||
// WEKNORA_PLUGIN_PYTHON, or python3 (python on Windows) from PATH.
|
||||
func PythonCommand() string {
|
||||
if c := strings.TrimSpace(os.Getenv("WEKNORA_PLUGIN_PYTHON")); c != "" {
|
||||
return c
|
||||
}
|
||||
if runtime.GOOS == "windows" {
|
||||
return "python"
|
||||
}
|
||||
return "python3"
|
||||
}
|
||||
|
||||
// EntryName is runtime.entry resolved for this machine: {os} and {arch}
|
||||
// filled in, and .exe added for Windows binaries.
|
||||
func EntryName(rt manifest.Runtime) string {
|
||||
rel := strings.NewReplacer("{os}", runtime.GOOS, "{arch}", runtime.GOARCH).Replace(rt.Entry)
|
||||
if rt.Kind == KindBinary && runtime.GOOS == "windows" && path.Ext(rel) == "" {
|
||||
rel += ".exe"
|
||||
}
|
||||
return rel
|
||||
}
|
||||
|
||||
// entryPath resolves runtime.entry for this machine inside the package.
|
||||
func entryPath(m *manifest.Manifest, dir string) (string, error) {
|
||||
rel := EntryName(m.Runtime)
|
||||
p, err := utils.SafeJoinUnderBase(dir, rel)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("runtime.entry %q: %w", m.Runtime.Entry, err)
|
||||
@@ -68,6 +98,9 @@ func entryPath(m *manifest.Manifest, dir string) (string, error) {
|
||||
if !info.Mode().IsRegular() {
|
||||
return "", fmt.Errorf("runtime.entry %s is not a file", rel)
|
||||
}
|
||||
if m.Runtime.Kind != KindBinary {
|
||||
return p, nil
|
||||
}
|
||||
// Packages are extracted without modes; the entry has to be executable.
|
||||
if err := os.Chmod(p, 0o755); err != nil {
|
||||
return "", err
|
||||
@@ -154,10 +187,10 @@ func randomToken() string {
|
||||
// childEnv is the plugin's whole environment. It is built from scratch so
|
||||
// WeKnora's own secrets (database, AES key, vendor keys) never reach plugin
|
||||
// code.
|
||||
func (p *process) childEnv(socket, token string) []string {
|
||||
func (p *process) childEnv(network, socket, token string) []string {
|
||||
env := []string{
|
||||
pluginapi.EnvSocket + "=" + socket,
|
||||
pluginapi.EnvNetwork + "=unix",
|
||||
pluginapi.EnvNetwork + "=" + network,
|
||||
pluginapi.EnvToken + "=" + token,
|
||||
pluginapi.EnvPluginID + "=" + p.spec.m.ID,
|
||||
pluginapi.EnvPluginVersion + "=" + p.spec.m.Version,
|
||||
@@ -177,18 +210,59 @@ func (p *process) childEnv(socket, token string) []string {
|
||||
env = append(env, k+"="+v)
|
||||
}
|
||||
}
|
||||
if p.spec.m.Runtime.Kind == KindPython {
|
||||
// Dependencies are vendored into the package; the extracted package
|
||||
// is shared, so no bytecode is written into it.
|
||||
env = append(env,
|
||||
"PYTHONPATH="+filepath.Join(p.spec.dir, "vendor")+string(os.PathListSeparator)+p.spec.dir,
|
||||
"PYTHONDONTWRITEBYTECODE=1",
|
||||
"PYTHONIOENCODING=utf-8",
|
||||
)
|
||||
}
|
||||
return env
|
||||
}
|
||||
|
||||
// listenAddress is where the plugin is told to listen: a private unix
|
||||
// socket, or a loopback port for Python on Windows, which has no unix
|
||||
// sockets.
|
||||
func (p *process) listenAddress() (network, address string) {
|
||||
if p.spec.m.Runtime.Kind == KindPython && runtime.GOOS == "windows" {
|
||||
return "tcp", "127.0.0.1:0"
|
||||
}
|
||||
return "unix", filepath.Join(p.sockDir, "p.sock")
|
||||
}
|
||||
|
||||
// loopback reports whether a TCP handshake names this machine.
|
||||
func loopback(hs pluginapi.Handshake) bool {
|
||||
if hs.Network != "tcp" {
|
||||
return true
|
||||
}
|
||||
host, _, err := net.SplitHostPort(hs.Address)
|
||||
ip := net.ParseIP(host)
|
||||
return err == nil && ip != nil && ip.IsLoopback()
|
||||
}
|
||||
|
||||
// command runs the entry: directly for binaries, with the interpreter for
|
||||
// python plugins (-s: without the user's site-packages, -u: unbuffered so
|
||||
// the handshake is not held back).
|
||||
func (p *process) command(entry string) *exec.Cmd {
|
||||
if p.spec.m.Runtime.Kind == KindPython {
|
||||
return exec.Command(PythonCommand(), "-s", "-u", entry)
|
||||
}
|
||||
return exec.Command(entry)
|
||||
}
|
||||
|
||||
// launch starts the child and waits for handshake, health and a manifest
|
||||
// that matches the installed package.
|
||||
func (p *process) launch(ctx context.Context, entry string) (*launched, error) {
|
||||
socket := filepath.Join(p.sockDir, "p.sock")
|
||||
_ = os.Remove(socket)
|
||||
network, socket := p.listenAddress()
|
||||
if network == "unix" {
|
||||
_ = os.Remove(socket)
|
||||
}
|
||||
token := randomToken()
|
||||
cmd := exec.Command(entry)
|
||||
cmd := p.command(entry)
|
||||
cmd.Dir = p.spec.dir
|
||||
cmd.Env = p.childEnv(socket, token)
|
||||
cmd.Env = p.childEnv(network, socket, token)
|
||||
configureChild(cmd)
|
||||
stdout, err := cmd.StdoutPipe()
|
||||
if err != nil {
|
||||
@@ -230,9 +304,9 @@ func (p *process) launch(ctx context.Context, entry string) (*launched, error) {
|
||||
kill()
|
||||
return nil, ctx.Err()
|
||||
}
|
||||
if hs.Network == "unix" && hs.Address != socket {
|
||||
if hs.Network != network || (network == "unix" && hs.Address != socket) || !loopback(hs) {
|
||||
kill()
|
||||
return nil, fmt.Errorf("plugin listens on %s, not the socket it was given", hs.Address)
|
||||
return nil, fmt.Errorf("plugin listens on %s %s, not where it was told", hs.Network, hs.Address)
|
||||
}
|
||||
c := client.ForHandshake(hs, client.Bearer(token))
|
||||
cctx, cancel := context.WithTimeout(ctx, healthTimeout)
|
||||
@@ -248,7 +322,7 @@ func (p *process) launch(ctx context.Context, entry string) (*launched, error) {
|
||||
kill()
|
||||
return nil, fmt.Errorf("read plugin manifest: %w", err)
|
||||
}
|
||||
if err := checkServedManifest(p.spec.m, m); err != nil {
|
||||
if err := CheckServedManifest(p.spec.m, m); err != nil {
|
||||
c.Close()
|
||||
kill()
|
||||
return nil, err
|
||||
@@ -256,14 +330,14 @@ func (p *process) launch(ctx context.Context, entry string) (*launched, error) {
|
||||
return &launched{cmd: cmd, client: c, exited: exited}, nil
|
||||
}
|
||||
|
||||
// checkServedManifest makes sure the process is the package that was
|
||||
// CheckServedManifest makes sure a running plugin is the package that was
|
||||
// installed and serves what the manifest promised.
|
||||
func checkServedManifest(want *manifest.Manifest, got *pluginapi.Manifest) error {
|
||||
func CheckServedManifest(want *manifest.Manifest, got *pluginapi.Manifest) error {
|
||||
if got.ID != want.ID || got.Version != want.Version {
|
||||
return fmt.Errorf("process reports %s@%s, the package is %s@%s", got.ID, got.Version, want.ID, want.Version)
|
||||
return fmt.Errorf("plugin reports %s@%s, the package is %s@%s", got.ID, got.Version, want.ID, want.Version)
|
||||
}
|
||||
if got.APIVersion != pluginapi.APIVersion {
|
||||
return fmt.Errorf("process speaks %q, this WeKnora speaks %q", got.APIVersion, pluginapi.APIVersion)
|
||||
return fmt.Errorf("plugin speaks %q, this WeKnora speaks %q", got.APIVersion, pluginapi.APIVersion)
|
||||
}
|
||||
for point, contribs := range want.Contributes {
|
||||
if info, ok := manifest.LookupPoint(point); !ok || info.Declarative {
|
||||
@@ -275,7 +349,7 @@ func checkServedManifest(want *manifest.Manifest, got *pluginapi.Manifest) error
|
||||
}
|
||||
for _, c := range contribs {
|
||||
if !served[c.ID] {
|
||||
return fmt.Errorf("plugin.yaml declares %s/%s but the process does not serve it", point, c.ID)
|
||||
return fmt.Errorf("plugin.yaml declares %s/%s but the plugin does not serve it", point, c.ID)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,130 @@
|
||||
package host
|
||||
|
||||
import (
|
||||
"context"
|
||||
"io/fs"
|
||||
"os"
|
||||
"os/exec"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/Tencent/WeKnora/internal/plugin/manifest"
|
||||
"github.com/Tencent/WeKnora/internal/plugin/reconcile"
|
||||
"github.com/Tencent/WeKnora/pluginsdk/pluginapi"
|
||||
)
|
||||
|
||||
// pythonEcho is the echo plugin written with the Python SDK.
|
||||
const pythonEcho = `import os
|
||||
from weknora_plugin import Plugin, SearchResult
|
||||
|
||||
plugin = Plugin(os.environ["WEKNORA_PLUGIN_ID"], os.environ["WEKNORA_PLUGIN_VERSION"])
|
||||
|
||||
|
||||
@plugin.web_search("echo")
|
||||
def search(call, q):
|
||||
if q.query == "env":
|
||||
return [SearchResult(title=",".join(sorted(os.environ)), url="env")]
|
||||
if q.query == "crash":
|
||||
os._exit(3)
|
||||
if q.query == "pid":
|
||||
return [SearchResult(title=str(os.getpid()), url="pid")]
|
||||
return [SearchResult(title=q.query, url="echo")]
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
plugin.serve()
|
||||
`
|
||||
|
||||
// installPython lays out an extracted python package: main.py and the SDK
|
||||
// vendored under vendor/, as a plugin author ships it.
|
||||
func installPython(t *testing.T) *reconcile.Loaded {
|
||||
t.Helper()
|
||||
if _, err := exec.LookPath(PythonCommand()); err != nil {
|
||||
t.Skipf("%s is not installed", PythonCommand())
|
||||
}
|
||||
dir := t.TempDir()
|
||||
sdk := filepath.Join("..", "..", "..", "pluginsdk", "python", "src", "weknora_plugin")
|
||||
err := filepath.WalkDir(sdk, func(p string, d fs.DirEntry, err error) error {
|
||||
if err != nil || d.IsDir() || filepath.Ext(p) != ".py" {
|
||||
return err
|
||||
}
|
||||
data, err := os.ReadFile(p)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
target := filepath.Join(dir, "vendor", "weknora_plugin", filepath.Base(p))
|
||||
if err := os.MkdirAll(filepath.Dir(target), 0o755); err != nil {
|
||||
return err
|
||||
}
|
||||
return os.WriteFile(target, data, 0o644)
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := os.WriteFile(filepath.Join(dir, "main.py"), []byte(pythonEcho), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
m := &manifest.Manifest{
|
||||
SchemaVersion: manifest.SchemaVersion, ID: "acme.echo", Version: "1.0.0", APIVersion: pluginapi.APIVersion,
|
||||
Name: manifest.Text("Echo", nil), Publisher: manifest.Publisher{ID: "acme"},
|
||||
Runtime: manifest.Runtime{Type: manifest.RuntimeHost, Kind: KindPython, Entry: "main.py"},
|
||||
Contributes: manifest.Contributions{
|
||||
manifest.PointWebSearch: {{ID: "echo", Name: manifest.Text("Echo", nil)}},
|
||||
},
|
||||
}
|
||||
return &reconcile.Loaded{Manifest: m, Dir: dir}
|
||||
}
|
||||
|
||||
func TestHostRunsAPythonPlugin(t *testing.T) {
|
||||
fastTimings(t)
|
||||
t.Setenv("DB_PASSWORD", "must-not-leak")
|
||||
ctx := context.Background()
|
||||
m := NewManager()
|
||||
rep := &reports{}
|
||||
m.SetReporter(rep)
|
||||
defer m.Close()
|
||||
|
||||
l := installPython(t)
|
||||
if err := m.Activate(ctx, l); err != nil {
|
||||
t.Fatalf("Activate: %v", err)
|
||||
}
|
||||
if got, err := search(t, m, "hello"); err != nil || got != "hello" {
|
||||
t.Fatalf("search = %q, %v", got, err)
|
||||
}
|
||||
env, err := search(t, m, "env")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if strings.Contains(env, "DB_PASSWORD") || !strings.Contains(env, "PYTHONPATH") ||
|
||||
!strings.Contains(env, pluginapi.EnvToken) {
|
||||
t.Fatalf("plugin environment = %s", env)
|
||||
}
|
||||
|
||||
pid1, _ := search(t, m, "pid")
|
||||
_, _ = search(t, m, "crash")
|
||||
waitFor(t, "restart", func() bool {
|
||||
pid, err := search(t, m, "pid")
|
||||
return err == nil && pid != pid1
|
||||
})
|
||||
if err := m.Deactivate(ctx, "acme.echo"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if entries, _ := os.ReadDir(l.Dir); len(entries) != 2 {
|
||||
t.Fatalf("the plugin wrote into its package: %v", entries)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEntryName(t *testing.T) {
|
||||
py := manifest.Runtime{Kind: KindPython, Entry: "main.py"}
|
||||
if got := EntryName(py); got != "main.py" {
|
||||
t.Fatalf("python entry = %s", got)
|
||||
}
|
||||
bin := manifest.Runtime{Kind: KindBinary, Entry: "bin/{os}-{arch}/x"}
|
||||
if got := EntryName(bin); strings.Contains(got, "{") {
|
||||
t.Fatalf("binary entry = %s", got)
|
||||
}
|
||||
if !Supported(KindPython) || Supported("node") {
|
||||
t.Fatal("this host runs binaries and python")
|
||||
}
|
||||
}
|
||||
@@ -6,13 +6,15 @@ package install
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"path"
|
||||
"os/exec"
|
||||
goruntime "runtime"
|
||||
"sort"
|
||||
"strings"
|
||||
@@ -21,6 +23,7 @@ import (
|
||||
"golang.org/x/mod/semver"
|
||||
|
||||
"github.com/Tencent/WeKnora/internal/logger"
|
||||
"github.com/Tencent/WeKnora/internal/plugin/host"
|
||||
"github.com/Tencent/WeKnora/internal/plugin/manifest"
|
||||
"github.com/Tencent/WeKnora/internal/plugin/pkg"
|
||||
"github.com/Tencent/WeKnora/internal/plugin/reconcile"
|
||||
@@ -47,23 +50,29 @@ func invalid(format string, args ...any) error {
|
||||
var supportedRuntimes = map[manifest.RuntimeType]bool{
|
||||
manifest.RuntimeDeclarative: true,
|
||||
manifest.RuntimeHost: true,
|
||||
manifest.RuntimeRemote: true,
|
||||
}
|
||||
|
||||
// checkHostRuntime makes sure this server can run a host plugin: a binary
|
||||
// built for its OS and architecture. Other kinds need a host image that
|
||||
// carries their interpreter.
|
||||
// built for its OS and architecture, or a Python entry and an interpreter.
|
||||
func checkHostRuntime(p *pkg.Package) error {
|
||||
rt := p.Manifest.Runtime
|
||||
if rt.Kind != "binary" {
|
||||
return invalid("runtime.kind %q is not supported yet; host plugins must be binaries", rt.Kind)
|
||||
}
|
||||
entry := strings.NewReplacer("{os}", goruntime.GOOS, "{arch}", goruntime.GOARCH).Replace(rt.Entry)
|
||||
if goruntime.GOOS == "windows" && path.Ext(entry) == "" {
|
||||
entry += ".exe"
|
||||
if !host.Supported(rt.Kind) {
|
||||
return invalid("runtime.kind %q is not supported yet; host plugins must be binaries or python", rt.Kind)
|
||||
}
|
||||
entry := host.EntryName(rt)
|
||||
if _, ok := p.ReadFile(entry); !ok {
|
||||
return invalid("the package has no build for this server (%s/%s): %s is missing",
|
||||
goruntime.GOOS, goruntime.GOARCH, entry)
|
||||
if rt.Kind == host.KindBinary {
|
||||
return invalid("the package has no build for this server (%s/%s): %s is missing",
|
||||
goruntime.GOOS, goruntime.GOARCH, entry)
|
||||
}
|
||||
return invalid("the package has no %s (runtime.entry)", entry)
|
||||
}
|
||||
if rt.Kind == host.KindPython {
|
||||
if _, err := exec.LookPath(host.PythonCommand()); err != nil {
|
||||
return invalid("python plugins need %s on this server; install it or set WEKNORA_PLUGIN_PYTHON",
|
||||
host.PythonCommand())
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -142,6 +151,9 @@ type View struct {
|
||||
Versions []types.PluginVersion `json:"versions"`
|
||||
// Node is the plugin's state on the node that served the request.
|
||||
Node *reconcile.Status `json:"node,omitempty"`
|
||||
// IssuedSecret is a remote plugin's signing secret, returned only by
|
||||
// the call that created it: the service needs it to check requests.
|
||||
IssuedSecret string `json:"issuedSecret,omitempty"`
|
||||
}
|
||||
|
||||
// Inspect opens a package and says what installing it would do.
|
||||
@@ -176,7 +188,7 @@ func (s *Service) open(data []byte) (*pkg.Package, error) {
|
||||
return nil, &InvalidError{Err: err}
|
||||
}
|
||||
if !supportedRuntimes[p.Manifest.Runtime.Type] {
|
||||
return nil, invalid("runtime %q is not supported yet; declarative and host plugins can be installed",
|
||||
return nil, invalid("runtime %q is not supported yet; declarative, host and remote plugins can be installed",
|
||||
p.Manifest.Runtime.Type)
|
||||
}
|
||||
if p.Manifest.Runtime.Type == manifest.RuntimeHost {
|
||||
@@ -203,6 +215,9 @@ type Request struct {
|
||||
// to the package the administrator reviewed.
|
||||
ExpectedDigest string
|
||||
UserID string
|
||||
// RemoteURL is where a remote plugin's service runs. An upgrade may
|
||||
// leave it empty to keep the registered one.
|
||||
RemoteURL string
|
||||
}
|
||||
|
||||
// Install stores the package as a version of its plugin and makes it the
|
||||
@@ -248,6 +263,23 @@ func (s *Service) Install(ctx context.Context, req Request) (*View, error) {
|
||||
if row == nil {
|
||||
row = &types.InstalledPlugin{ID: m.ID, DesiredState: types.PluginStateEnabled, CreatedBy: req.UserID}
|
||||
}
|
||||
var issued string
|
||||
if m.Runtime.Type == manifest.RuntimeRemote {
|
||||
if req.RemoteURL != "" {
|
||||
if err := checkRemoteURL(req.RemoteURL); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
row.RemoteURL = strings.TrimSuffix(req.RemoteURL, "/")
|
||||
}
|
||||
if row.RemoteURL == "" {
|
||||
return nil, invalid("a remote plugin needs the URL of its service")
|
||||
}
|
||||
if row.RemoteSecret == "" {
|
||||
if issued, err = s.issueSecret(row); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
}
|
||||
source, _ := json.Marshal(req.Source)
|
||||
perms, _ := json.Marshal(m.Permissions)
|
||||
row.Source = types.JSON(source)
|
||||
@@ -258,7 +290,94 @@ func (s *Service) Install(ctx context.Context, req Request) (*View, error) {
|
||||
return nil, err
|
||||
}
|
||||
logger.Infof(ctx, "[plugin] %s installed %s %s (%s)", req.UserID, m.ID, m.Version, p.Digest)
|
||||
return s.apply(ctx, m.ID)
|
||||
v, err := s.apply(ctx, m.ID)
|
||||
if v != nil {
|
||||
v.IssuedSecret = issued
|
||||
}
|
||||
return v, err
|
||||
}
|
||||
|
||||
// checkRemoteURL accepts an http(s) service URL that passes the SSRF rules;
|
||||
// a service on a private network needs its host in SSRF_WHITELIST.
|
||||
func checkRemoteURL(raw string) error {
|
||||
u, err := url.Parse(raw)
|
||||
if err != nil || (u.Scheme != "https" && u.Scheme != "http") || u.Host == "" {
|
||||
return invalid("service URL must be an http(s) URL")
|
||||
}
|
||||
if u.User != nil || u.RawQuery != "" || u.Fragment != "" {
|
||||
return invalid("service URL must not carry credentials, a query or a fragment")
|
||||
}
|
||||
if err := utils.ValidateURLForSSRF(raw); err != nil {
|
||||
return invalid("service URL is not allowed (private hosts must be in SSRF_WHITELIST): %v", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// issueSecret gives a remote plugin a new signing secret, stored sealed,
|
||||
// and returns it in the clear for the administrator to hand to the service.
|
||||
func (s *Service) issueSecret(row *types.InstalledPlugin) (string, error) {
|
||||
b := make([]byte, 32)
|
||||
if _, err := rand.Read(b); err != nil {
|
||||
return "", err
|
||||
}
|
||||
secret := hex.EncodeToString(b)
|
||||
sealed, err := utils.EncryptAESGCM(secret, utils.GetAESKey())
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
row.RemoteSecret = sealed
|
||||
return secret, nil
|
||||
}
|
||||
|
||||
// SetRemoteURL moves a remote plugin to another service URL.
|
||||
func (s *Service) SetRemoteURL(ctx context.Context, id, rawURL string) (*View, error) {
|
||||
row, err := s.remoteRow(ctx, id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := checkRemoteURL(rawURL); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
row.RemoteURL = strings.TrimSuffix(rawURL, "/")
|
||||
if err := s.repo.SavePlugin(ctx, row); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return s.apply(ctx, id)
|
||||
}
|
||||
|
||||
// RotateSecret replaces a remote plugin's signing secret. Calls fail until
|
||||
// the service is given the new one, which only this response shows.
|
||||
func (s *Service) RotateSecret(ctx context.Context, id string) (*View, error) {
|
||||
row, err := s.remoteRow(ctx, id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
secret, err := s.issueSecret(row)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := s.repo.SavePlugin(ctx, row); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
v, err := s.apply(ctx, id)
|
||||
if v != nil {
|
||||
v.IssuedSecret = secret
|
||||
}
|
||||
return v, err
|
||||
}
|
||||
|
||||
func (s *Service) remoteRow(ctx context.Context, id string) (*types.InstalledPlugin, error) {
|
||||
row, err := s.repo.GetPlugin(ctx, id)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if row == nil {
|
||||
return nil, ErrNotInstalled
|
||||
}
|
||||
if row.Runtime != string(manifest.RuntimeRemote) {
|
||||
return nil, invalid("%s is not a remote plugin", id)
|
||||
}
|
||||
return row, nil
|
||||
}
|
||||
|
||||
// FetchURL downloads a package over HTTP(S), refusing private addresses.
|
||||
|
||||
@@ -3,6 +3,7 @@ package install
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"os"
|
||||
goruntime "runtime"
|
||||
"strings"
|
||||
"testing"
|
||||
@@ -11,6 +12,7 @@ import (
|
||||
"github.com/Tencent/WeKnora/internal/plugin/reconcile"
|
||||
"github.com/Tencent/WeKnora/internal/plugin/registry"
|
||||
"github.com/Tencent/WeKnora/internal/types"
|
||||
"github.com/Tencent/WeKnora/internal/utils"
|
||||
)
|
||||
|
||||
func newService(t *testing.T) (*Service, *plugintest.MemRepo, *plugintest.MemStore, *registry.Registry) {
|
||||
@@ -156,8 +158,104 @@ func TestInstallHostPlugins(t *testing.T) {
|
||||
!strings.Contains(err.Error(), "no build for this server") {
|
||||
t.Fatalf("want a missing build error, got %v", err)
|
||||
}
|
||||
if _, err := s.Inspect(ctx, hostPackage(t, "python", here)); !isInvalid(err) ||
|
||||
!strings.Contains(err.Error(), "must be binaries") {
|
||||
if _, err := s.Inspect(ctx, hostPackage(t, "node", here)); !isInvalid(err) ||
|
||||
!strings.Contains(err.Error(), "binaries or python") {
|
||||
t.Fatalf("want an unsupported kind error, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestInstallPythonPlugins(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
s, _, _, _ := newService(t)
|
||||
pkg := func(files map[string]string) []byte {
|
||||
files["plugin.yaml"] = "schemaVersion: 1\nid: acme.py\nversion: 1.0.0\napiVersion: weknora.plugin/v1\n" +
|
||||
"name: { en-US: ACME Py }\npublisher: { id: acme }\n" +
|
||||
"runtime: { type: host, kind: python, entry: main.py }\n" +
|
||||
"contributes:\n webSearch:\n - { id: search, name: ACME Search }\n"
|
||||
return plugintest.Zip(t, files)
|
||||
}
|
||||
if _, err := s.Inspect(ctx, pkg(map[string]string{})); !isInvalid(err) ||
|
||||
!strings.Contains(err.Error(), "no main.py") {
|
||||
t.Fatalf("want a missing entry error, got %v", err)
|
||||
}
|
||||
t.Setenv("WEKNORA_PLUGIN_PYTHON", "no-such-python-here")
|
||||
if _, err := s.Inspect(ctx, pkg(map[string]string{"main.py": "print()"})); !isInvalid(err) ||
|
||||
!strings.Contains(err.Error(), "no-such-python-here") {
|
||||
t.Fatalf("want a missing interpreter error, got %v", err)
|
||||
}
|
||||
t.Setenv("WEKNORA_PLUGIN_PYTHON", os.Args[0]) // any executable will do
|
||||
if _, err := s.Inspect(ctx, pkg(map[string]string{"main.py": "print()"})); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
func remotePackage(t *testing.T, version string) []byte {
|
||||
return plugintest.Zip(t, map[string]string{
|
||||
"plugin.yaml": "schemaVersion: 1\nid: acme.remote\nversion: " + version +
|
||||
"\napiVersion: weknora.plugin/v1\n" +
|
||||
"name: { en-US: ACME Remote }\npublisher: { id: acme }\nruntime: { type: remote }\n" +
|
||||
"contributes:\n webSearch:\n - { id: search, name: ACME Search }\n",
|
||||
})
|
||||
}
|
||||
|
||||
func TestInstallRemotePlugins(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
utils.SetSSRFWhitelistFromRaw("plugins.example.com")
|
||||
t.Cleanup(func() { utils.SetSSRFWhitelistFromRaw("") })
|
||||
t.Setenv("SYSTEM_AES_KEY", strings.Repeat("k", 32))
|
||||
s, repo, _, _ := newService(t)
|
||||
|
||||
if _, err := s.Install(ctx, Request{Data: remotePackage(t, "1.0.0")}); !isInvalid(err) ||
|
||||
!strings.Contains(err.Error(), "URL of its service") {
|
||||
t.Fatalf("want a missing URL error, got %v", err)
|
||||
}
|
||||
refused := []string{"http://127.0.0.1:9000", "ftp://plugins.example.com", "https://u:p@plugins.example.com"}
|
||||
for _, u := range refused {
|
||||
if _, err := s.Install(ctx, Request{Data: remotePackage(t, "1.0.0"), RemoteURL: u}); !isInvalid(err) {
|
||||
t.Errorf("RemoteURL %s = %v, want refusal", u, err)
|
||||
}
|
||||
}
|
||||
|
||||
view, err := s.Install(ctx, Request{
|
||||
Data: remotePackage(t, "1.0.0"), RemoteURL: "https://plugins.example.com/acme/",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(view.IssuedSecret) != 64 || view.RemoteURL != "https://plugins.example.com/acme" {
|
||||
t.Fatalf("view = %+v", view)
|
||||
}
|
||||
row, _ := repo.GetPlugin(ctx, "acme.remote")
|
||||
if plain, _ := utils.DecryptStoredSecret(row.RemoteSecret); !strings.HasPrefix(row.RemoteSecret, utils.EncPrefix) ||
|
||||
plain != view.IssuedSecret {
|
||||
t.Fatalf("stored secret %q does not seal the issued one", row.RemoteSecret)
|
||||
}
|
||||
first := view.IssuedSecret
|
||||
|
||||
view, err = s.Install(ctx, Request{Data: remotePackage(t, "1.1.0")})
|
||||
if err != nil || view.IssuedSecret != "" || view.RemoteURL != "https://plugins.example.com/acme" {
|
||||
t.Fatalf("upgrade must keep the URL and secret: %+v, %v", view, err)
|
||||
}
|
||||
|
||||
view, err = s.RotateSecret(ctx, "acme.remote")
|
||||
if err != nil || len(view.IssuedSecret) != 64 || view.IssuedSecret == first {
|
||||
t.Fatalf("rotate = %+v, %v", view, err)
|
||||
}
|
||||
if view, err = s.SetRemoteURL(ctx, "acme.remote", "https://plugins.example.com/v2"); err != nil ||
|
||||
view.RemoteURL != "https://plugins.example.com/v2" || view.IssuedSecret != "" {
|
||||
t.Fatalf("move = %+v, %v", view, err)
|
||||
}
|
||||
if _, err := s.SetRemoteURL(ctx, "acme.remote", "http://10.0.0.1"); !isInvalid(err) {
|
||||
t.Fatalf("private URL = %v", err)
|
||||
}
|
||||
|
||||
if _, err := s.Install(ctx, Request{Data: plugintest.KitPackage(t, "1.0.0")}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := s.RotateSecret(ctx, "acme.kit"); !isInvalid(err) {
|
||||
t.Fatalf("rotating a declarative plugin = %v", err)
|
||||
}
|
||||
if _, err := s.RotateSecret(ctx, "acme.none"); !errors.Is(err, ErrNotInstalled) {
|
||||
t.Fatalf("rotating a missing plugin = %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -8,6 +8,8 @@ package reconcile
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
@@ -37,6 +39,9 @@ type Loaded struct {
|
||||
// Dir is the package extracted on local disk, for consumers that read
|
||||
// files by path (skills).
|
||||
Dir string
|
||||
// Installed is the plugin's row: where a remote plugin runs and its
|
||||
// sealed secret.
|
||||
Installed types.InstalledPlugin
|
||||
}
|
||||
|
||||
// Activator wires one domain to plugin contributions: it registers what a
|
||||
@@ -101,14 +106,40 @@ type Reconciler struct {
|
||||
instanceID string
|
||||
interval time.Duration
|
||||
|
||||
mu sync.Mutex // serializes passes
|
||||
loaded map[string]*Loaded
|
||||
digests map[string]string // plugin ID → loaded digest
|
||||
mu sync.Mutex // serializes passes
|
||||
loaded map[string]*Loaded
|
||||
digests map[string]string // plugin ID → loaded digest and runtime target
|
||||
// retries holds plugins whose activation failed: a process that would
|
||||
// not start, a remote service that was down. They are tried again with
|
||||
// backoff until they load or change.
|
||||
retries map[string]retry
|
||||
now func() time.Time
|
||||
statusMu sync.RWMutex
|
||||
status map[string]Status
|
||||
runOnce sync.Once
|
||||
}
|
||||
|
||||
// retry is when a failed activation is tried again.
|
||||
type retry struct {
|
||||
key string // the load key that failed
|
||||
attempts int
|
||||
next time.Time
|
||||
}
|
||||
|
||||
// Backoff between attempts to activate a plugin that failed.
|
||||
const (
|
||||
retryFloor = 30 * time.Second
|
||||
retryCeiling = 10 * time.Minute
|
||||
)
|
||||
|
||||
func retryDelay(attempts int) time.Duration {
|
||||
d := retryFloor
|
||||
for i := 1; i < attempts && d < retryCeiling; i++ {
|
||||
d *= 2
|
||||
}
|
||||
return min(d, retryCeiling)
|
||||
}
|
||||
|
||||
// Options configures a Reconciler.
|
||||
type Options struct {
|
||||
Repo interfaces.PluginRepository
|
||||
@@ -131,7 +162,8 @@ func New(o Options) *Reconciler {
|
||||
return &Reconciler{
|
||||
repo: o.Repo, store: o.Store, registry: o.Registry, cacheDir: o.CacheDir, rdb: o.Redis,
|
||||
activators: o.Activators, instanceID: uuid.NewString(), interval: o.Interval,
|
||||
loaded: map[string]*Loaded{}, digests: map[string]string{}, status: map[string]Status{},
|
||||
loaded: map[string]*Loaded{}, digests: map[string]string{}, retries: map[string]retry{},
|
||||
status: map[string]Status{}, now: time.Now,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -165,7 +197,7 @@ func (r *Reconciler) Reconcile(ctx context.Context) error {
|
||||
r.setStatus(row.ID, Status{Version: row.ActiveVersion, State: StateFailed, Error: err.Error()})
|
||||
}
|
||||
}
|
||||
for id := range r.digests {
|
||||
for id := range r.loaded {
|
||||
if !desired[id] {
|
||||
r.unload(ctx, id)
|
||||
}
|
||||
@@ -183,6 +215,16 @@ func (r *Reconciler) Reconcile(ctx context.Context) error {
|
||||
return errors.Join(errs...)
|
||||
}
|
||||
|
||||
// runtimeTarget identifies where a plugin runs beyond its package: a new
|
||||
// remote URL or secret reloads the plugin like a new version would.
|
||||
func runtimeTarget(row types.InstalledPlugin) string {
|
||||
if row.RemoteURL == "" && row.RemoteSecret == "" {
|
||||
return ""
|
||||
}
|
||||
sum := sha256.Sum256([]byte(row.RemoteURL + "\x00" + row.RemoteSecret))
|
||||
return hex.EncodeToString(sum[:8])
|
||||
}
|
||||
|
||||
// ensure loads the active version of one plugin unless it already is.
|
||||
func (r *Reconciler) ensure(ctx context.Context, row types.InstalledPlugin) error {
|
||||
v, err := r.repo.GetVersion(ctx, row.ID, row.ActiveVersion)
|
||||
@@ -192,9 +234,13 @@ func (r *Reconciler) ensure(ctx context.Context, row types.InstalledPlugin) erro
|
||||
if v == nil {
|
||||
return fmt.Errorf("version %s is not stored", row.ActiveVersion)
|
||||
}
|
||||
if r.digests[row.ID] == v.Digest {
|
||||
loadKey := v.Digest + "|" + runtimeTarget(row)
|
||||
if r.digests[row.ID] == loadKey {
|
||||
return nil
|
||||
}
|
||||
if rt, ok := r.retries[row.ID]; ok && rt.key == loadKey && r.now().Before(rt.next) {
|
||||
return nil // still failed; its status says why
|
||||
}
|
||||
data, err := r.store.Get(ctx, v.PackageURI)
|
||||
if err != nil {
|
||||
return err
|
||||
@@ -216,7 +262,7 @@ func (r *Reconciler) ensure(ctx context.Context, row types.InstalledPlugin) erro
|
||||
if err := r.registry.Replace(p.Manifest); err != nil {
|
||||
return err
|
||||
}
|
||||
l := &Loaded{Manifest: p.Manifest, Package: p, Dir: dir}
|
||||
l := &Loaded{Manifest: p.Manifest, Package: p, Dir: dir, Installed: row}
|
||||
var errs []error
|
||||
for _, a := range r.activators {
|
||||
_, inPlace := a.(InPlaceActivator)
|
||||
@@ -230,10 +276,19 @@ func (r *Reconciler) ensure(ctx context.Context, row types.InstalledPlugin) erro
|
||||
}
|
||||
}
|
||||
r.loaded[row.ID] = l
|
||||
r.digests[row.ID] = v.Digest
|
||||
if err := errors.Join(errs...); err != nil {
|
||||
rt := r.retries[row.ID]
|
||||
if rt.key != loadKey {
|
||||
rt = retry{key: loadKey}
|
||||
}
|
||||
rt.attempts++
|
||||
rt.next = r.now().Add(retryDelay(rt.attempts))
|
||||
r.retries[row.ID] = rt
|
||||
delete(r.digests, row.ID)
|
||||
return err
|
||||
}
|
||||
r.digests[row.ID] = loadKey
|
||||
delete(r.retries, row.ID)
|
||||
logger.Infof(ctx, "[plugin] loaded %s %s", row.ID, row.ActiveVersion)
|
||||
r.setStatus(row.ID, Status{Version: row.ActiveVersion, State: StateReady})
|
||||
return nil
|
||||
@@ -252,6 +307,7 @@ func (r *Reconciler) unload(ctx context.Context, id string) {
|
||||
}
|
||||
delete(r.loaded, id)
|
||||
delete(r.digests, id)
|
||||
delete(r.retries, id)
|
||||
r.statusMu.Lock()
|
||||
delete(r.status, id)
|
||||
r.statusMu.Unlock()
|
||||
|
||||
@@ -7,6 +7,7 @@ import (
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/Tencent/WeKnora/internal/plugin/manifest"
|
||||
"github.com/Tencent/WeKnora/internal/plugin/plugintest"
|
||||
@@ -128,3 +129,100 @@ func TestActivatorFailureMarksPluginFailed(t *testing.T) {
|
||||
t.Fatalf("status = %+v", s)
|
||||
}
|
||||
}
|
||||
|
||||
// A plugin whose activation failed (a service that was down) is tried again
|
||||
// with backoff, and loads once the cause is gone.
|
||||
func TestFailedActivationIsRetriedWithBackoff(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
repo, store, act := plugintest.NewMemRepo(), &plugintest.MemStore{}, &recorder{fail: true}
|
||||
r := New(Options{
|
||||
Repo: repo, Store: store, Registry: registry.New(), CacheDir: t.TempDir(), Activators: []Activator{act},
|
||||
})
|
||||
now := time.Unix(1000, 0)
|
||||
r.now = func() time.Time { return now }
|
||||
plugintest.Install(t, repo, store, plugintest.KitPackage(t, "1.0.0"), types.PluginStateEnabled)
|
||||
|
||||
if err := r.Reconcile(ctx); err == nil {
|
||||
t.Fatal("want the activation error")
|
||||
}
|
||||
now = now.Add(retryFloor - time.Second)
|
||||
_ = r.Reconcile(ctx)
|
||||
if len(act.calls) != 1 {
|
||||
t.Fatalf("retried before the backoff: %v", act.calls)
|
||||
}
|
||||
now = now.Add(time.Second)
|
||||
_ = r.Reconcile(ctx)
|
||||
if len(act.calls) != 3 { // deactivate the half-loaded plugin, activate again
|
||||
t.Fatalf("calls = %v", act.calls)
|
||||
}
|
||||
now = now.Add(retryFloor)
|
||||
_ = r.Reconcile(ctx)
|
||||
if len(act.calls) != 3 {
|
||||
t.Fatalf("the second retry should wait twice as long: %v", act.calls)
|
||||
}
|
||||
|
||||
act.fail = false
|
||||
now = now.Add(retryFloor)
|
||||
if err := r.Reconcile(ctx); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if s, _ := r.Status("acme.kit"); s.State != StateReady {
|
||||
t.Fatalf("status = %+v", s)
|
||||
}
|
||||
_ = r.Reconcile(ctx)
|
||||
if len(act.calls) != 5 {
|
||||
t.Fatalf("a loaded plugin should stay loaded: %v", act.calls)
|
||||
}
|
||||
|
||||
// A failed plugin that is disabled is unloaded all the same.
|
||||
act.fail = true
|
||||
plugintest.Install(t, repo, store, plugintest.KitPackage(t, "1.1.0"), types.PluginStateEnabled)
|
||||
_ = r.Reconcile(ctx)
|
||||
row, _ := repo.GetPlugin(ctx, "acme.kit")
|
||||
row.DesiredState = types.PluginStateDisabled
|
||||
_ = repo.SavePlugin(ctx, row)
|
||||
_ = r.Reconcile(ctx)
|
||||
if len(r.Loaded()) != 0 {
|
||||
t.Fatal("a disabled plugin must be unloaded even if it failed")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRetryDelay(t *testing.T) {
|
||||
if retryDelay(1) != retryFloor || retryDelay(2) != 2*retryFloor || retryDelay(50) != retryCeiling {
|
||||
t.Fatalf("delays = %s %s %s", retryDelay(1), retryDelay(2), retryDelay(50))
|
||||
}
|
||||
}
|
||||
|
||||
// A remote plugin moved to another URL, or given a new secret, is loaded
|
||||
// again so its runtime picks the change up.
|
||||
func TestReconcileReloadsOnRuntimeTargetChange(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
repo, store, reg, act := plugintest.NewMemRepo(), &plugintest.MemStore{}, registry.New(), &recorder{}
|
||||
r := New(Options{Repo: repo, Store: store, Registry: reg, CacheDir: t.TempDir(), Activators: []Activator{act}})
|
||||
|
||||
plugintest.Install(t, repo, store, plugintest.KitPackage(t, "1.0.0"), types.PluginStateEnabled)
|
||||
row, _ := repo.GetPlugin(ctx, "acme.kit")
|
||||
row.RemoteURL, row.RemoteSecret = "https://a.example.com", "s1"
|
||||
_ = repo.SavePlugin(ctx, row)
|
||||
if err := r.Reconcile(ctx); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if got := r.Loaded()[0].Installed.RemoteURL; got != "https://a.example.com" {
|
||||
t.Fatalf("Loaded carries URL %q", got)
|
||||
}
|
||||
_ = r.Reconcile(ctx)
|
||||
|
||||
row.RemoteURL = "https://b.example.com"
|
||||
_ = repo.SavePlugin(ctx, row)
|
||||
_ = r.Reconcile(ctx)
|
||||
row.RemoteSecret = "s2"
|
||||
_ = repo.SavePlugin(ctx, row)
|
||||
_ = r.Reconcile(ctx)
|
||||
|
||||
if len(act.calls) != 5 || act.calls[1] != "deactivate acme.kit" {
|
||||
t.Fatalf("calls = %v, want a reload per change", act.calls)
|
||||
}
|
||||
if got := r.Loaded()[0].Installed.RemoteSecret; got != "s2" {
|
||||
t.Fatalf("Loaded carries secret %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,248 @@
|
||||
// Package remote reaches remote-runtime plugins: HTTP services an
|
||||
// administrator runs somewhere else and registers by URL. Every request is
|
||||
// signed with the plugin's shared secret, so the service can tell WeKnora
|
||||
// from anyone else who can reach it.
|
||||
package remote
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/Tencent/WeKnora/internal/logger"
|
||||
"github.com/Tencent/WeKnora/internal/plugin/host"
|
||||
"github.com/Tencent/WeKnora/internal/plugin/manifest"
|
||||
"github.com/Tencent/WeKnora/internal/plugin/reconcile"
|
||||
"github.com/Tencent/WeKnora/internal/utils"
|
||||
"github.com/Tencent/WeKnora/pluginsdk/client"
|
||||
"github.com/Tencent/WeKnora/pluginsdk/pluginapi"
|
||||
)
|
||||
|
||||
const (
|
||||
healthInterval = 15 * time.Second
|
||||
checkTimeout = 10 * time.Second
|
||||
)
|
||||
|
||||
// Manager keeps a client for every loaded remote plugin and health-checks
|
||||
// it. It is a reconcile.Activator and, like the embedded host, must come
|
||||
// before the activators that route calls to the plugins.
|
||||
type Manager struct {
|
||||
newClient func(url string, secret []byte) *client.Client
|
||||
interval time.Duration
|
||||
|
||||
mu sync.Mutex
|
||||
reporter host.StateReporter
|
||||
endpoints map[string]*endpoint
|
||||
}
|
||||
|
||||
// endpoint is one registered remote plugin.
|
||||
type endpoint struct {
|
||||
m *manifest.Manifest
|
||||
c *client.Client
|
||||
cancel context.CancelFunc
|
||||
done chan struct{}
|
||||
|
||||
mu sync.Mutex
|
||||
// err is why the service cannot be called right now: unreachable,
|
||||
// or serving something other than the installed package.
|
||||
err error
|
||||
}
|
||||
|
||||
// NewManager creates a Manager whose clients refuse private addresses
|
||||
// unless SSRF_WHITELIST allows them.
|
||||
func NewManager() *Manager {
|
||||
cfg := utils.DefaultSSRFSafeHTTPClientConfig()
|
||||
cfg.Timeout = 0 // calls are bounded by their context; syncs stream
|
||||
cfg.SameOriginRedirectsOnly = true
|
||||
httpClient := utils.NewSSRFSafeHTTPClient(cfg)
|
||||
return &Manager{
|
||||
newClient: func(url string, secret []byte) *client.Client {
|
||||
return client.New(url, httpClient, client.Signed(secret))
|
||||
},
|
||||
interval: healthInterval,
|
||||
endpoints: map[string]*endpoint{},
|
||||
}
|
||||
}
|
||||
|
||||
// SetReporter wires health changes to the reconciler.
|
||||
func (m *Manager) SetReporter(r host.StateReporter) {
|
||||
m.mu.Lock()
|
||||
m.reporter = r
|
||||
m.mu.Unlock()
|
||||
}
|
||||
|
||||
// Name implements reconcile.Activator.
|
||||
func (m *Manager) Name() string { return "remote" }
|
||||
|
||||
// ActivatesInPlace implements reconcile.InPlaceActivator: a new URL or
|
||||
// version takes over from the old endpoint without a gap.
|
||||
func (m *Manager) ActivatesInPlace() {}
|
||||
|
||||
// Activate checks that the registered service is up and serves the
|
||||
// installed package, then keeps health-checking it. Other runtimes are
|
||||
// ignored.
|
||||
func (m *Manager) Activate(ctx context.Context, l *reconcile.Loaded) error {
|
||||
if l.Manifest.Runtime.Type != manifest.RuntimeRemote {
|
||||
return nil
|
||||
}
|
||||
url := l.Installed.RemoteURL
|
||||
if url == "" {
|
||||
return errors.New("no service URL is registered for this remote plugin")
|
||||
}
|
||||
if err := utils.ValidateURLForSSRF(url); err != nil {
|
||||
return fmt.Errorf("service URL is not allowed: %w", err)
|
||||
}
|
||||
secret, err := utils.DecryptStoredSecret(l.Installed.RemoteSecret)
|
||||
if err != nil {
|
||||
return fmt.Errorf("decrypt the plugin secret: %w", err)
|
||||
}
|
||||
if secret == "" {
|
||||
return errors.New("the remote plugin has no signing secret")
|
||||
}
|
||||
c := m.newClient(url, []byte(secret))
|
||||
cctx, cancel := context.WithTimeout(ctx, checkTimeout)
|
||||
err = verify(cctx, c, l.Manifest)
|
||||
cancel()
|
||||
if err != nil {
|
||||
c.Close()
|
||||
return err
|
||||
}
|
||||
|
||||
id := l.Manifest.ID
|
||||
wctx, stop := context.WithCancel(context.Background())
|
||||
e := &endpoint{m: l.Manifest, c: c, cancel: stop, done: make(chan struct{})}
|
||||
m.mu.Lock()
|
||||
old := m.endpoints[id]
|
||||
m.endpoints[id] = e
|
||||
m.mu.Unlock()
|
||||
if old != nil {
|
||||
old.close()
|
||||
}
|
||||
go m.watch(wctx, id, e)
|
||||
logger.Infof(ctx, "[plugin] remote %s %s at %s", id, l.Manifest.Version, url)
|
||||
return nil
|
||||
}
|
||||
|
||||
// verify checks that the service is healthy and is the installed package.
|
||||
func verify(ctx context.Context, c *client.Client, want *manifest.Manifest) error {
|
||||
if err := c.Health(ctx); err != nil {
|
||||
return fmt.Errorf("plugin service is not healthy: %w", err)
|
||||
}
|
||||
got, err := c.Manifest(ctx)
|
||||
if err != nil {
|
||||
return fmt.Errorf("read plugin manifest: %w", err)
|
||||
}
|
||||
return host.CheckServedManifest(want, got)
|
||||
}
|
||||
|
||||
// watch health-checks the service until the endpoint is replaced or
|
||||
// removed. After an outage it checks the manifest again: the service may
|
||||
// have been redeployed with another version.
|
||||
func (m *Manager) watch(ctx context.Context, id string, e *endpoint) {
|
||||
defer close(e.done)
|
||||
t := time.NewTicker(m.interval)
|
||||
defer t.Stop()
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case <-t.C:
|
||||
}
|
||||
cctx, cancel := context.WithTimeout(ctx, checkTimeout)
|
||||
var err error
|
||||
if e.failing() {
|
||||
err = verify(cctx, e.c, e.m)
|
||||
} else if err = e.c.Health(cctx); err != nil {
|
||||
err = fmt.Errorf("plugin service is not healthy: %w", err)
|
||||
}
|
||||
cancel()
|
||||
if ctx.Err() != nil {
|
||||
return
|
||||
}
|
||||
if e.setErr(err) {
|
||||
m.report(id, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (m *Manager) report(id string, err error) {
|
||||
m.mu.Lock()
|
||||
r := m.reporter
|
||||
m.mu.Unlock()
|
||||
if r != nil {
|
||||
r.ReportRuntime(id, err == nil, err)
|
||||
}
|
||||
}
|
||||
|
||||
func (e *endpoint) failing() bool {
|
||||
e.mu.Lock()
|
||||
defer e.mu.Unlock()
|
||||
return e.err != nil
|
||||
}
|
||||
|
||||
// setErr records the latest check and says whether health changed.
|
||||
func (e *endpoint) setErr(err error) bool {
|
||||
e.mu.Lock()
|
||||
defer e.mu.Unlock()
|
||||
changed := (e.err == nil) != (err == nil)
|
||||
e.err = err
|
||||
return changed
|
||||
}
|
||||
|
||||
func (e *endpoint) close() {
|
||||
e.cancel()
|
||||
<-e.done
|
||||
e.c.Close()
|
||||
}
|
||||
|
||||
// Deactivate implements reconcile.Activator.
|
||||
func (m *Manager) Deactivate(ctx context.Context, pluginID string) error {
|
||||
m.mu.Lock()
|
||||
e := m.endpoints[pluginID]
|
||||
delete(m.endpoints, pluginID)
|
||||
m.mu.Unlock()
|
||||
if e != nil {
|
||||
e.close()
|
||||
logger.Infof(ctx, "[plugin] remote %s unregistered", pluginID)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// Client returns the client of a registered remote plugin. The error is a
|
||||
// retryable pluginapi unavailable error while the service fails its checks.
|
||||
func (m *Manager) Client(pluginID string) (*client.Client, error) {
|
||||
m.mu.Lock()
|
||||
e := m.endpoints[pluginID]
|
||||
m.mu.Unlock()
|
||||
if e == nil {
|
||||
return nil, pluginapi.Errorf(pluginapi.CodeUnavailable, "remote plugin %s is not registered", pluginID)
|
||||
}
|
||||
e.mu.Lock()
|
||||
err := e.err
|
||||
e.mu.Unlock()
|
||||
if err != nil {
|
||||
return nil, pluginapi.Errorf(pluginapi.CodeUnavailable, "%v", err)
|
||||
}
|
||||
return e.c, nil
|
||||
}
|
||||
|
||||
// Owns says whether a plugin is registered here.
|
||||
func (m *Manager) Owns(pluginID string) bool {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
_, ok := m.endpoints[pluginID]
|
||||
return ok
|
||||
}
|
||||
|
||||
// Close unregisters every plugin, for shutdown.
|
||||
func (m *Manager) Close() {
|
||||
m.mu.Lock()
|
||||
eps := m.endpoints
|
||||
m.endpoints = map[string]*endpoint{}
|
||||
m.mu.Unlock()
|
||||
for _, e := range eps {
|
||||
e.close()
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,196 @@
|
||||
package remote
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"errors"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/Tencent/WeKnora/internal/plugin/manifest"
|
||||
"github.com/Tencent/WeKnora/internal/plugin/reconcile"
|
||||
"github.com/Tencent/WeKnora/internal/types"
|
||||
"github.com/Tencent/WeKnora/internal/utils"
|
||||
"github.com/Tencent/WeKnora/pluginsdk"
|
||||
"github.com/Tencent/WeKnora/pluginsdk/pluginapi"
|
||||
)
|
||||
|
||||
// service is a remote plugin that checks WeKnora's signature and can be
|
||||
// taken down.
|
||||
type service struct {
|
||||
*httptest.Server
|
||||
down atomic.Bool
|
||||
}
|
||||
|
||||
func newService(t *testing.T, version, secret string) *service {
|
||||
t.Helper()
|
||||
p := pluginsdk.New(pluginsdk.Info{ID: "acme.search", Version: version})
|
||||
p.WebSearch("web", pluginsdk.WebSearchFunc(
|
||||
func(context.Context, *pluginsdk.Call, pluginapi.SearchInput) (*pluginapi.SearchOutput, error) {
|
||||
return &pluginapi.SearchOutput{}, nil
|
||||
}))
|
||||
h := p.Handler()
|
||||
s := &service{}
|
||||
s.Server = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if s.down.Load() {
|
||||
http.Error(w, "down", http.StatusBadGateway)
|
||||
return
|
||||
}
|
||||
body, _ := io.ReadAll(r.Body)
|
||||
err := pluginapi.VerifySignature([]byte(secret), r.Header.Get(pluginapi.TimestampHeader),
|
||||
r.Header.Get(pluginapi.SignatureHeader), body, time.Now())
|
||||
if err != nil {
|
||||
http.Error(w, err.Error(), http.StatusUnauthorized)
|
||||
return
|
||||
}
|
||||
r.Body = io.NopCloser(bytes.NewReader(body))
|
||||
h.ServeHTTP(w, r)
|
||||
}))
|
||||
t.Cleanup(s.Close)
|
||||
return s
|
||||
}
|
||||
|
||||
type reports struct {
|
||||
mu sync.Mutex
|
||||
got []bool
|
||||
}
|
||||
|
||||
func (r *reports) ReportRuntime(_ string, healthy bool, _ error) {
|
||||
r.mu.Lock()
|
||||
r.got = append(r.got, healthy)
|
||||
r.mu.Unlock()
|
||||
}
|
||||
|
||||
func (r *reports) seen() []bool {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
return append([]bool(nil), r.got...)
|
||||
}
|
||||
|
||||
func loaded(version, url, secret string) *reconcile.Loaded {
|
||||
return &reconcile.Loaded{
|
||||
Manifest: &manifest.Manifest{
|
||||
ID: "acme.search", Version: version, APIVersion: manifest.ExtensionAPIVersion,
|
||||
Runtime: manifest.Runtime{Type: manifest.RuntimeRemote},
|
||||
Contributes: map[manifest.Point][]manifest.Contribution{
|
||||
manifest.PointWebSearch: {{ID: "web"}},
|
||||
},
|
||||
},
|
||||
Installed: types.InstalledPlugin{ID: "acme.search", RemoteURL: url, RemoteSecret: secret},
|
||||
}
|
||||
}
|
||||
|
||||
func allowLoopback(t *testing.T) {
|
||||
utils.SetSSRFWhitelistFromRaw("127.0.0.1")
|
||||
t.Cleanup(func() { utils.SetSSRFWhitelistFromRaw("") })
|
||||
}
|
||||
|
||||
func TestActivateVerifiesAndSignsCalls(t *testing.T) {
|
||||
allowLoopback(t)
|
||||
ctx := context.Background()
|
||||
svc := newService(t, "1.0.0", "s3cret")
|
||||
m := NewManager()
|
||||
defer m.Close()
|
||||
|
||||
if err := m.Activate(ctx, loaded("1.0.0", svc.URL, "wrong")); err == nil ||
|
||||
!strings.Contains(err.Error(), "not healthy") {
|
||||
t.Fatalf("wrong secret = %v", err)
|
||||
}
|
||||
if err := m.Activate(ctx, loaded("2.0.0", svc.URL, "s3cret")); err == nil ||
|
||||
!strings.Contains(err.Error(), "the package is acme.search@2.0.0") {
|
||||
t.Fatalf("other version = %v", err)
|
||||
}
|
||||
if err := m.Activate(ctx, loaded("1.0.0", "", "s3cret")); err == nil {
|
||||
t.Fatal("no URL should fail")
|
||||
}
|
||||
if err := m.Activate(ctx, loaded("1.0.0", svc.URL, "s3cret")); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
c, err := m.Client("acme.search")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
var out pluginapi.SearchOutput
|
||||
if err := c.Call(ctx, pluginapi.SearchPath("web"), pluginapi.Envelope{}, pluginapi.SearchInput{Query: "q"},
|
||||
&out); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
if err := m.Deactivate(ctx, "acme.search"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := m.Client("acme.search"); err == nil {
|
||||
t.Fatal("a removed plugin should have no client")
|
||||
}
|
||||
}
|
||||
|
||||
func TestActivateRefusesPrivateAddresses(t *testing.T) {
|
||||
svc := newService(t, "1.0.0", "s3cret")
|
||||
m := NewManager()
|
||||
defer m.Close()
|
||||
err := m.Activate(context.Background(), loaded("1.0.0", svc.URL, "s3cret"))
|
||||
if err == nil || !strings.Contains(err.Error(), "not allowed") {
|
||||
t.Fatalf("loopback without whitelist = %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSealedSecret(t *testing.T) {
|
||||
allowLoopback(t)
|
||||
t.Setenv("SYSTEM_AES_KEY", strings.Repeat("k", 32))
|
||||
sealed, err := utils.EncryptAESGCM("s3cret", utils.GetAESKey())
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
svc := newService(t, "1.0.0", "s3cret")
|
||||
m := NewManager()
|
||||
defer m.Close()
|
||||
if err := m.Activate(context.Background(), loaded("1.0.0", svc.URL, sealed)); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestHealthChecksReportOutages(t *testing.T) {
|
||||
allowLoopback(t)
|
||||
svc := newService(t, "1.0.0", "s3cret")
|
||||
m := NewManager()
|
||||
m.interval = 10 * time.Millisecond
|
||||
rep := &reports{}
|
||||
m.SetReporter(rep)
|
||||
defer m.Close()
|
||||
if err := m.Activate(context.Background(), loaded("1.0.0", svc.URL, "s3cret")); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
svc.down.Store(true)
|
||||
waitFor(t, func() bool { return len(rep.seen()) == 1 })
|
||||
var pe *pluginapi.Error
|
||||
if _, err := m.Client("acme.search"); !errors.As(err, &pe) || !pe.Retryable {
|
||||
t.Fatalf("client while down = %v", err)
|
||||
}
|
||||
|
||||
svc.down.Store(false)
|
||||
waitFor(t, func() bool { return len(rep.seen()) == 2 })
|
||||
if got := rep.seen(); got[0] || !got[1] {
|
||||
t.Fatalf("reports = %v, want degraded then healthy", got)
|
||||
}
|
||||
if _, err := m.Client("acme.search"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
func waitFor(t *testing.T, cond func() bool) {
|
||||
t.Helper()
|
||||
deadline := time.Now().Add(5 * time.Second)
|
||||
for !cond() {
|
||||
if time.Now().After(deadline) {
|
||||
t.Fatal("timed out")
|
||||
}
|
||||
time.Sleep(5 * time.Millisecond)
|
||||
}
|
||||
}
|
||||
@@ -39,6 +39,8 @@ func RegisterPluginAdminRoutes(r *gin.RouterGroup, h *handler.PluginAdminHandler
|
||||
plugins.DELETE("/:id", h.UninstallPlugin)
|
||||
plugins.PUT("/:id/enabled", h.SetInstalledPluginEnabled)
|
||||
plugins.PUT("/:id/active-version", h.ActivatePluginVersion)
|
||||
plugins.PUT("/:id/remote-url", h.SetPluginRemoteURL)
|
||||
plugins.POST("/:id/secret/rotate", h.RotatePluginSecret)
|
||||
plugins.GET("/:id/config", h.GetPluginSystemConfig)
|
||||
plugins.PUT("/:id/config", h.UpdatePluginSystemConfig)
|
||||
}
|
||||
|
||||
@@ -22,8 +22,12 @@ type InstalledPlugin struct {
|
||||
DesiredState string `json:"desired_state" gorm:"type:varchar(16)"`
|
||||
Runtime string `json:"runtime" gorm:"type:varchar(32)"`
|
||||
// GrantedPerms is the manifest permissions an administrator accepted.
|
||||
GrantedPerms JSON `json:"granted_perms" gorm:"type:json"`
|
||||
SystemConfig JSON `json:"system_config,omitempty" gorm:"type:json"`
|
||||
GrantedPerms JSON `json:"granted_perms" gorm:"type:json"`
|
||||
SystemConfig JSON `json:"system_config,omitempty" gorm:"type:json"`
|
||||
// RemoteURL is where a remote plugin is served; RemoteSecret (sealed)
|
||||
// signs the calls to it. Other runtimes leave both empty.
|
||||
RemoteURL string `json:"remote_url,omitempty" gorm:"type:varchar(1024)"`
|
||||
RemoteSecret string `json:"-" gorm:"type:text"`
|
||||
CreatedBy string `json:"created_by" gorm:"type:varchar(36)"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
|
||||
@@ -0,0 +1,2 @@
|
||||
ALTER TABLE plugins DROP COLUMN remote_secret;
|
||||
ALTER TABLE plugins DROP COLUMN remote_url;
|
||||
@@ -0,0 +1,3 @@
|
||||
-- Remote plugin URL and signing secret (versioned 000118).
|
||||
ALTER TABLE plugins ADD COLUMN remote_url VARCHAR(1024) NOT NULL DEFAULT '';
|
||||
ALTER TABLE plugins ADD COLUMN remote_secret TEXT NOT NULL DEFAULT '';
|
||||
@@ -0,0 +1,2 @@
|
||||
ALTER TABLE plugins DROP COLUMN IF EXISTS remote_secret;
|
||||
ALTER TABLE plugins DROP COLUMN IF EXISTS remote_url;
|
||||
@@ -0,0 +1,7 @@
|
||||
-- Migration 000118: where a remote plugin runs and the secret its calls are
|
||||
-- signed with. Only runtime.type remote plugins set them.
|
||||
DO $$ BEGIN RAISE NOTICE '[Migration 000118] Adding remote plugin columns'; END $$;
|
||||
|
||||
ALTER TABLE plugins ADD COLUMN IF NOT EXISTS remote_url VARCHAR(1024) NOT NULL DEFAULT '';
|
||||
-- Encrypted (enc:v1) HMAC secret shared with the plugin service.
|
||||
ALTER TABLE plugins ADD COLUMN IF NOT EXISTS remote_secret TEXT NOT NULL DEFAULT '';
|
||||
+42
-2
@@ -10,6 +10,9 @@ connectors and document parsers that run as their own process. The module has no
|
||||
| `client` | Call a plugin, as WeKnora does. |
|
||||
| `conformance`, `cmd/weknora-plugin-conformance` | Check that a plugin speaks the protocol. |
|
||||
|
||||
Writing Python? [`python/`](python/README.md) is the same SDK for Python
|
||||
3.9+, with only the standard library.
|
||||
|
||||
## A connector in 30 lines
|
||||
|
||||
```go
|
||||
@@ -119,6 +122,11 @@ A few rules the manifest and host enforce:
|
||||
|
||||
- **Binaries.** Build one per platform into `bin/<os>-<arch>/`; WeKnora
|
||||
runs the build for its own OS and architecture.
|
||||
- **Python.** A plugin can instead be Python source:
|
||||
- Declare `runtime: { type: host, kind: python, entry: main.py }`.
|
||||
- WeKnora runs the entry with its own `python3` (or
|
||||
`WEKNORA_PLUGIN_PYTHON`).
|
||||
- `vendor/` and the package root are on `PYTHONPATH`.
|
||||
- **Environment.** The process gets no environment from WeKnora beyond the
|
||||
`WEKNORA_PLUGIN_*` variables.
|
||||
- **Outbound traffic.** It goes through the host's egress proxy, which
|
||||
@@ -128,8 +136,11 @@ A few rules the manifest and host enforce:
|
||||
(`acme.notes/notes`); for connectors and web search that must stay within
|
||||
50 characters.
|
||||
|
||||
See `examples/plugins/rss` (a connector) and `examples/plugins/subtitles` (a
|
||||
parser using the Host API) for complete plugins with their `package.sh`.
|
||||
Complete plugins with their `package.sh`:
|
||||
|
||||
- `examples/plugins/rss`: a connector.
|
||||
- `examples/plugins/subtitles`: a parser using the Host API.
|
||||
- `examples/plugins/notebooks`: a parser written in Python.
|
||||
|
||||
## Testing
|
||||
|
||||
@@ -144,3 +155,32 @@ In Go tests, run `conformance.Run` against `httptest.NewServer(p.Handler())`.
|
||||
In WeKnora: **System administration → Plugin management → Install plugin**.
|
||||
Upload the `.wkp` or give its URL, review what it adds and reaches, and
|
||||
install. Each workspace then enables it under **Settings → Plugins**.
|
||||
|
||||
## Running as a remote service
|
||||
|
||||
A plugin can also run as a service you deploy yourself, for instance in its
|
||||
own container or on another team's cluster. Its package then carries only
|
||||
`plugin.yaml` (plus schemas):
|
||||
|
||||
```yaml
|
||||
runtime: { type: remote }
|
||||
```
|
||||
|
||||
1. **Install.** Give the service URL when installing. WeKnora checks the
|
||||
service's `/v1/manifest` against the package: same ID, version and
|
||||
contributions.
|
||||
2. **Keep the secret.** Installing shows a signing secret once. Start the
|
||||
service with it as `WEKNORA_PLUGIN_SECRET` (and `WEKNORA_PLUGIN_ADDR`,
|
||||
default `:8080`). The SDK rejects requests without a valid signature.
|
||||
3. **Private hosts.** A service on a private network must be listed in
|
||||
WeKnora's `SSRF_WHITELIST`.
|
||||
4. **Host API.** Remote plugins get a Host API token only when
|
||||
`WEKNORA_PLUGIN_HOST_API_URL` tells WeKnora its address as the service
|
||||
sees it.
|
||||
|
||||
The plugin detail page changes the URL and rotates the secret. After a
|
||||
rotation, calls fail until the service has the new secret.
|
||||
|
||||
Upgrade the service and the package together. While the service reports a
|
||||
version other than the active package, calls are refused and the plugin
|
||||
shows as degraded.
|
||||
|
||||
@@ -0,0 +1,175 @@
|
||||
package conformance_test
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"bytes"
|
||||
"context"
|
||||
"net"
|
||||
"net/http"
|
||||
"os"
|
||||
"os/exec"
|
||||
"path/filepath"
|
||||
"runtime"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/Tencent/WeKnora/pluginsdk/client"
|
||||
"github.com/Tencent/WeKnora/pluginsdk/conformance"
|
||||
"github.com/Tencent/WeKnora/pluginsdk/pluginapi"
|
||||
)
|
||||
|
||||
// The Python SDK speaks the same protocol: its test plugin must pass the
|
||||
// suite in both modes, and a Go client must read its answers.
|
||||
|
||||
func python(t *testing.T) string {
|
||||
t.Helper()
|
||||
names := []string{"python3", "python"}
|
||||
if runtime.GOOS == "windows" {
|
||||
names = []string{"python", "python3"} // python3 may be the Store alias
|
||||
}
|
||||
for _, name := range names {
|
||||
if p, err := exec.LookPath(name); err == nil {
|
||||
return p
|
||||
}
|
||||
}
|
||||
t.Skip("python3 is not installed")
|
||||
return ""
|
||||
}
|
||||
|
||||
var fixture = filepath.Join("..", "python", "tests", "fixture.py")
|
||||
|
||||
func startPython(t *testing.T, env ...string) *exec.Cmd {
|
||||
t.Helper()
|
||||
cmd := exec.Command(python(t), fixture)
|
||||
cmd.Env = append([]string{"PATH=" + os.Getenv("PATH")}, env...)
|
||||
cmd.Stderr = os.Stderr
|
||||
return cmd
|
||||
}
|
||||
|
||||
func stopPython(t *testing.T, cmd *exec.Cmd) {
|
||||
t.Cleanup(func() {
|
||||
if runtime.GOOS == "windows" {
|
||||
_ = cmd.Process.Kill()
|
||||
} else {
|
||||
_ = cmd.Process.Signal(os.Interrupt)
|
||||
}
|
||||
done := make(chan struct{})
|
||||
go func() { _ = cmd.Wait(); close(done) }()
|
||||
select {
|
||||
case <-done:
|
||||
case <-time.After(10 * time.Second):
|
||||
_ = cmd.Process.Kill()
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func report(t *testing.T, rep conformance.Report) {
|
||||
t.Helper()
|
||||
for _, r := range rep.Results {
|
||||
if !r.Passed {
|
||||
t.Errorf("%s: %s", r.Name, r.Detail)
|
||||
}
|
||||
}
|
||||
if len(rep.Results) < 13 {
|
||||
t.Fatalf("only %d checks ran", len(rep.Results))
|
||||
}
|
||||
}
|
||||
|
||||
func TestPythonPluginConformsInHostMode(t *testing.T) {
|
||||
if runtime.GOOS == "windows" {
|
||||
t.Skip("the Python SDK serves host mode on unix sockets")
|
||||
}
|
||||
dir, err := os.MkdirTemp("", "wkpy")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Cleanup(func() { _ = os.RemoveAll(dir) })
|
||||
sock := filepath.Join(dir, "p.sock")
|
||||
cmd := startPython(t, pluginapi.EnvSocket+"="+sock, pluginapi.EnvToken+"=tok")
|
||||
stdout, _ := cmd.StdoutPipe()
|
||||
if err := cmd.Start(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
stopPython(t, cmd)
|
||||
line, err := bufio.NewReader(stdout).ReadString('\n')
|
||||
if err != nil {
|
||||
t.Fatalf("no handshake: %v", err)
|
||||
}
|
||||
hs, ok, err := pluginapi.ParseHandshake(line)
|
||||
if !ok || err != nil || hs.Address != sock {
|
||||
t.Fatalf("handshake %q: %v", line, err)
|
||||
}
|
||||
c := client.ForHandshake(hs, client.Bearer("tok"))
|
||||
report(t, conformance.Run(context.Background(), conformance.Target{
|
||||
Client: c, Unauthenticated: client.ForHandshake(hs, nil),
|
||||
Raw: func(ctx context.Context, path string, body []byte) (*http.Response, error) {
|
||||
tr := &http.Transport{DialContext: func(ctx context.Context, _, _ string) (net.Conn, error) {
|
||||
return (&net.Dialer{}).DialContext(ctx, "unix", sock)
|
||||
}}
|
||||
req, _ := http.NewRequestWithContext(ctx, http.MethodPost, "http://plugin"+path, bytes.NewReader(body))
|
||||
req.Header.Set("Authorization", "Bearer tok")
|
||||
return (&http.Client{Transport: tr}).Do(req)
|
||||
},
|
||||
}))
|
||||
|
||||
// A Go client reads the Python stream and errors like a Go plugin's.
|
||||
var items int
|
||||
cursor, err := c.Stream(context.Background(), pluginapi.ConnectorFetchPath("notes"), pluginapi.Envelope{},
|
||||
pluginapi.FetchInput{Mode: pluginapi.FetchFull}, func(ev pluginapi.Event) error {
|
||||
if ev.Type == pluginapi.EventItem {
|
||||
items++
|
||||
}
|
||||
return nil
|
||||
})
|
||||
if err != nil || items != 3 || string(cursor) != `{"state":{"after":3}}` {
|
||||
t.Fatalf("fetch = %d items, cursor %s, %v", items, cursor, err)
|
||||
}
|
||||
var out pluginapi.ParseOutput
|
||||
err = c.Call(context.Background(), pluginapi.ParsePath("upper"), pluginapi.Envelope{},
|
||||
pluginapi.ParseInput{FileType: "txt", Content: []byte("hi")}, &out)
|
||||
if err != nil || len(out.Images) != 1 || string(out.Images[0].Data) != "\x89PNG" {
|
||||
t.Fatalf("parse = %+v, %v", out, err)
|
||||
}
|
||||
err = c.Call(context.Background(), pluginapi.SearchPath("echo"), pluginapi.Envelope{},
|
||||
pluginapi.SearchInput{Query: "down"}, nil)
|
||||
if pe, ok := pluginapi.AsError(err); !ok || pe.Code != pluginapi.CodeUnavailable || !pe.Retryable {
|
||||
t.Fatalf("search error = %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPythonPluginConformsInRemoteMode(t *testing.T) {
|
||||
ln, err := net.Listen("tcp", "127.0.0.1:0")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
addr := ln.Addr().String()
|
||||
_ = ln.Close()
|
||||
cmd := startPython(t, pluginapi.EnvAddr+"="+addr, pluginapi.EnvSecret+"=s3cret")
|
||||
if err := cmd.Start(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
stopPython(t, cmd)
|
||||
base := "http://" + addr
|
||||
c := client.New(base, nil, client.Signed("s3cret"))
|
||||
deadline := time.Now().Add(30 * time.Second)
|
||||
for {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
|
||||
err := c.Health(ctx)
|
||||
cancel()
|
||||
if err == nil {
|
||||
break
|
||||
}
|
||||
if time.Now().After(deadline) {
|
||||
t.Fatalf("the Python plugin did not come up: %v", err)
|
||||
}
|
||||
time.Sleep(50 * time.Millisecond)
|
||||
}
|
||||
report(t, conformance.Run(context.Background(), conformance.Target{
|
||||
Client: c, Unauthenticated: client.New(base, nil, nil),
|
||||
Raw: func(ctx context.Context, path string, body []byte) (*http.Response, error) {
|
||||
req, _ := http.NewRequestWithContext(ctx, http.MethodPost, base+path, bytes.NewReader(body))
|
||||
_ = client.Signed("s3cret").Apply(req, body)
|
||||
return http.DefaultClient.Do(req)
|
||||
},
|
||||
}))
|
||||
}
|
||||
@@ -0,0 +1,3 @@
|
||||
build/
|
||||
dist/
|
||||
*.egg-info/
|
||||
@@ -0,0 +1,133 @@
|
||||
# weknora-plugin
|
||||
|
||||
Build [WeKnora](https://github.com/Tencent/WeKnora) code plugins in Python.
|
||||
The package speaks the same extension protocol as the Go SDK. It uses only
|
||||
the standard library and supports Python 3.9 and later.
|
||||
|
||||
```python
|
||||
from weknora_plugin import Plugin, SearchResult
|
||||
|
||||
plugin = Plugin("acme.search", "1.0.0") # must match plugin.yaml
|
||||
|
||||
|
||||
@plugin.web_search("web")
|
||||
def search(call, q):
|
||||
# call.tenant_id, call.locale, call.system / call.tenant / call.instance
|
||||
return [SearchResult(title=q.query, url="https://example.com")]
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
plugin.serve()
|
||||
```
|
||||
|
||||
## Contributions
|
||||
|
||||
**Web search.**
|
||||
|
||||
```python
|
||||
@plugin.web_search(id)
|
||||
def search(call, q: SearchInput) -> list[SearchResult]: ...
|
||||
```
|
||||
|
||||
**Parser.** Return a `ParseOutput`, or just the Markdown as a `str`.
|
||||
|
||||
```python
|
||||
@plugin.parser(id)
|
||||
def parse(call, doc: ParseInput) -> ParseOutput: ...
|
||||
```
|
||||
|
||||
- `doc.content` is the file as bytes.
|
||||
- Images the Markdown references as `` go in `images` with a
|
||||
matching `original_ref`.
|
||||
|
||||
**Connector.** A class, registered with `@plugin.connector(id)` (or pass
|
||||
an instance):
|
||||
|
||||
```python
|
||||
@plugin.connector("notes")
|
||||
class Notes:
|
||||
def validate(self, call, cfg: ConnectorConfig) -> None: ...
|
||||
def list_resources(self, call, cfg, parent_id: str) -> list[Resource]: ...
|
||||
def fetch(self, call, cfg, inp: FetchInput, stream: Stream) -> Cursor | None:
|
||||
for note in changed_since(inp.cursor):
|
||||
stream.item(FetchedItem(external_id=note.id, title=note.title, content=note.markdown))
|
||||
stream.checkpoint(Cursor(state={"after": last_id})) # at page boundaries
|
||||
return Cursor(state={"after": last_id})
|
||||
# optional: resolve_ancestors(self, call, cfg, resource_ids) -> list[str]
|
||||
```
|
||||
|
||||
**Configuration check.**
|
||||
|
||||
```python
|
||||
@plugin.config_validator
|
||||
def validate(call) -> None: ...
|
||||
```
|
||||
|
||||
## Errors
|
||||
|
||||
Raise `PluginError(ErrorCode.X, "message")` to choose what WeKnora sees:
|
||||
|
||||
- `unauthorized`: credentials stopped working.
|
||||
- `rate_limited` (with `retry_after=`): back off.
|
||||
- `unavailable`: retry later.
|
||||
- `invalid_config(message, {"settings.url": "required"})`: point at form
|
||||
fields.
|
||||
|
||||
Any other exception becomes `internal`, and its traceback goes to stderr.
|
||||
|
||||
## Calling back into WeKnora
|
||||
|
||||
A plugin granted Host API scopes (`permissions.hostApi` in plugin.yaml) gets
|
||||
a client per call from `call.host()`. It returns `None` without a grant.
|
||||
|
||||
```python
|
||||
host = call.host()
|
||||
host.kv_put("cursor", {"page": 3}, ttl=3600)
|
||||
host.kv_get("cursor", default={})
|
||||
```
|
||||
|
||||
The token behind it lasts a few minutes: use it within the call.
|
||||
|
||||
## Running
|
||||
|
||||
`plugin.serve()` picks its mode from the environment.
|
||||
|
||||
- **Host mode.** The WeKnora plugin host starts the plugin with
|
||||
`WEKNORA_PLUGIN_SOCKET` and `WEKNORA_PLUGIN_TOKEN`. The plugin listens
|
||||
on that socket and prints the handshake on stdout. Keep stdout free of
|
||||
anything else, and log to stderr (`plugin.logger`).
|
||||
- **Remote mode.** Run on its own, it listens on `WEKNORA_PLUGIN_ADDR`
|
||||
(default `:8080`). It accepts only requests signed with
|
||||
`WEKNORA_PLUGIN_SECRET`, the secret shown when the plugin was registered.
|
||||
|
||||
On SIGTERM it stops accepting calls and lets the ones in flight finish.
|
||||
Each thread handles one call. `call.deadline` says when WeKnora stops
|
||||
waiting.
|
||||
|
||||
## Packaging for the plugin host
|
||||
|
||||
```yaml
|
||||
runtime: { type: host, kind: python, entry: main.py }
|
||||
```
|
||||
|
||||
WeKnora runs the entry with its own `python3`. It puts the package's
|
||||
`vendor/` directory and the package root on `PYTHONPATH`. Vendor the SDK
|
||||
and any dependencies:
|
||||
|
||||
```bash
|
||||
pip install --target vendor weknora-plugin # or copy src/weknora_plugin
|
||||
```
|
||||
|
||||
Pure-Python dependencies work everywhere. Ones with native code must match
|
||||
the server's platform. For a complete plugin, see
|
||||
`examples/plugins/notebooks` (a Jupyter notebook parser).
|
||||
|
||||
## Testing
|
||||
|
||||
`plugin.test_server()` serves the plugin on a loopback port without
|
||||
authentication, for unit tests. To check the protocol end to end, run the
|
||||
conformance suite from the Go SDK against a running plugin:
|
||||
|
||||
```bash
|
||||
go run github.com/Tencent/WeKnora/pluginsdk/cmd/weknora-plugin-conformance -url http://localhost:8080 -secret $SECRET
|
||||
```
|
||||
@@ -0,0 +1,23 @@
|
||||
[build-system]
|
||||
requires = ["setuptools>=64"]
|
||||
build-backend = "setuptools.build_meta"
|
||||
|
||||
[project]
|
||||
name = "weknora-plugin"
|
||||
version = "0.1.0"
|
||||
description = "Build WeKnora code plugins in Python"
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.9"
|
||||
license = { text = "MIT" }
|
||||
authors = [{ name = "WeKnora" }]
|
||||
classifiers = [
|
||||
"Programming Language :: Python :: 3",
|
||||
"Operating System :: OS Independent",
|
||||
]
|
||||
dependencies = []
|
||||
|
||||
[project.urls]
|
||||
Homepage = "https://github.com/Tencent/WeKnora/tree/main/pluginsdk/python"
|
||||
|
||||
[tool.setuptools.packages.find]
|
||||
where = ["src"]
|
||||
@@ -0,0 +1,67 @@
|
||||
"""Build WeKnora code plugins in Python.
|
||||
|
||||
from weknora_plugin import Plugin, SearchResult
|
||||
|
||||
plugin = Plugin("acme.search", "1.0.0")
|
||||
|
||||
@plugin.web_search("web")
|
||||
def search(call, query):
|
||||
return [SearchResult(title="WeKnora", url="https://github.com/Tencent/WeKnora")]
|
||||
|
||||
if __name__ == "__main__":
|
||||
plugin.serve()
|
||||
|
||||
Started by the WeKnora plugin host, the plugin listens where the host says
|
||||
and accepts only the host's token. Run on its own (a remote plugin), it
|
||||
listens on WEKNORA_PLUGIN_ADDR and checks request signatures with
|
||||
WEKNORA_PLUGIN_SECRET. The package uses only the standard library.
|
||||
"""
|
||||
|
||||
from .host import Host
|
||||
from .plugin import Call, Plugin, Stream, StreamClosed
|
||||
from .protocol import API_VERSION, PROTOCOL_VERSION, ErrorCode, PluginError, invalid_config
|
||||
from .types import (
|
||||
FETCH_FULL,
|
||||
FETCH_INCREMENTAL,
|
||||
ConnectorConfig,
|
||||
Cursor,
|
||||
FetchedItem,
|
||||
FetchInput,
|
||||
KVEntry,
|
||||
KVList,
|
||||
ParsedImage,
|
||||
ParseInput,
|
||||
ParseOutput,
|
||||
Resource,
|
||||
SearchInput,
|
||||
SearchResult,
|
||||
)
|
||||
|
||||
__version__ = "0.1.0"
|
||||
|
||||
__all__ = [
|
||||
"API_VERSION",
|
||||
"PROTOCOL_VERSION",
|
||||
"FETCH_FULL",
|
||||
"FETCH_INCREMENTAL",
|
||||
"Call",
|
||||
"ConnectorConfig",
|
||||
"Cursor",
|
||||
"ErrorCode",
|
||||
"FetchInput",
|
||||
"FetchedItem",
|
||||
"Host",
|
||||
"KVEntry",
|
||||
"KVList",
|
||||
"ParseInput",
|
||||
"ParseOutput",
|
||||
"ParsedImage",
|
||||
"Plugin",
|
||||
"PluginError",
|
||||
"Resource",
|
||||
"SearchInput",
|
||||
"SearchResult",
|
||||
"Stream",
|
||||
"StreamClosed",
|
||||
"invalid_config",
|
||||
]
|
||||
@@ -0,0 +1,95 @@
|
||||
"""The Host API client: how a plugin calls back into WeKnora during a call."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import urllib.error
|
||||
import urllib.parse
|
||||
import urllib.request
|
||||
from datetime import timedelta
|
||||
from typing import Any, Optional, Tuple, Union
|
||||
|
||||
from .protocol import HOST_KV_LIST_PATH, HOST_KV_PATH, ErrorCode, PluginError
|
||||
from .types import KVEntry, KVList, from_wire
|
||||
|
||||
|
||||
class Host:
|
||||
"""Reaches the Host API with the short-lived token of one call. Use it
|
||||
within the call; the token expires in minutes."""
|
||||
|
||||
def __init__(self, url: str, token: str, timeout: float = 30) -> None:
|
||||
self._base = url.rstrip("/")
|
||||
self._token = token
|
||||
self._timeout = timeout
|
||||
|
||||
def _do(self, method: str, path: str, query: Optional[dict] = None, body: Any = None) -> Any:
|
||||
url = self._base + path
|
||||
if query:
|
||||
url += "?" + urllib.parse.urlencode(query)
|
||||
data = None
|
||||
headers = {"Authorization": "Bearer " + self._token}
|
||||
if body is not None:
|
||||
data = json.dumps(body).encode()
|
||||
headers["Content-Type"] = "application/json"
|
||||
req = urllib.request.Request(url, data=data, method=method, headers=headers)
|
||||
try:
|
||||
with urllib.request.urlopen(req, timeout=self._timeout) as resp:
|
||||
raw = resp.read()
|
||||
except urllib.error.HTTPError as e:
|
||||
raw = e.read()
|
||||
try:
|
||||
err = json.loads(raw)["error"]
|
||||
if err.get("code"):
|
||||
raise PluginError.from_wire(err) from None
|
||||
except (ValueError, KeyError, TypeError):
|
||||
pass
|
||||
raise PluginError(ErrorCode.INTERNAL, f"host api answered HTTP {e.code}") from None
|
||||
except (urllib.error.URLError, OSError) as e:
|
||||
raise PluginError(ErrorCode.UNAVAILABLE, f"host api: {e}") from None
|
||||
return json.loads(raw) if raw else None
|
||||
|
||||
def kv_entry(self, key: str) -> Optional[KVEntry]:
|
||||
"""Reads a key with its metadata, or None when it does not exist."""
|
||||
try:
|
||||
return from_wire(KVEntry, self._do("GET", HOST_KV_PATH, {"key": key}))
|
||||
except PluginError as e:
|
||||
if e.code == ErrorCode.NOT_FOUND:
|
||||
return None
|
||||
raise
|
||||
|
||||
def kv_get(self, key: str, default: Any = None) -> Any:
|
||||
"""Reads a key's value, or default when it does not exist."""
|
||||
entry = self.kv_entry(key)
|
||||
return default if entry is None else entry.value
|
||||
|
||||
def kv_put(self, key: str, value: Any, ttl: Union[float, timedelta] = 0) -> None:
|
||||
"""Stores a JSON value; ttl (seconds or a timedelta) 0 keeps it
|
||||
until deleted."""
|
||||
if isinstance(ttl, timedelta):
|
||||
ttl = ttl.total_seconds()
|
||||
body: dict = {"key": key, "value": value}
|
||||
if ttl:
|
||||
body["ttlSeconds"] = int(ttl)
|
||||
self._do("PUT", HOST_KV_PATH, body=body)
|
||||
|
||||
def kv_delete(self, key: str) -> None:
|
||||
"""Removes a key; a missing key is not an error."""
|
||||
self._do("DELETE", HOST_KV_PATH, {"key": key})
|
||||
|
||||
def kv_list(self, prefix: str = "", after: str = "", limit: int = 0) -> KVList:
|
||||
"""A page of keys with prefix, in key order after the given key;
|
||||
pass the result's next as after for the following page."""
|
||||
q = {"prefix": prefix, "after": after}
|
||||
if limit > 0:
|
||||
q["limit"] = str(limit)
|
||||
return from_wire(KVList, self._do("GET", HOST_KV_LIST_PATH, q))
|
||||
|
||||
def kv_items(self, prefix: str = "") -> "Tuple[KVEntry, ...]":
|
||||
"""Every entry with prefix, following pages."""
|
||||
out, after = [], ""
|
||||
while True:
|
||||
page = self.kv_list(prefix, after)
|
||||
out.extend(page.entries)
|
||||
if not page.next:
|
||||
return tuple(out)
|
||||
after = page.next
|
||||
@@ -0,0 +1,559 @@
|
||||
"""Plugin: register contributions, then serve the protocol."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import hmac
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import re
|
||||
import signal
|
||||
import socket
|
||||
import socketserver
|
||||
import sys
|
||||
import threading
|
||||
import time
|
||||
import traceback
|
||||
from datetime import datetime, timezone
|
||||
from http.server import BaseHTTPRequestHandler, HTTPServer
|
||||
from typing import Any, Callable, Dict, Iterator, List, Optional, Tuple
|
||||
|
||||
from . import protocol as p
|
||||
from .host import Host
|
||||
from .protocol import ErrorCode, PluginError
|
||||
from .types import (
|
||||
ConnectorConfig,
|
||||
Cursor,
|
||||
FetchInput,
|
||||
FetchedItem,
|
||||
ParseInput,
|
||||
ParseOutput,
|
||||
SearchInput,
|
||||
SearchResult,
|
||||
from_wire,
|
||||
parse_time,
|
||||
to_wire,
|
||||
)
|
||||
|
||||
|
||||
class Call:
|
||||
"""One request's caller and configuration. Configuration secrets are
|
||||
already decrypted; do not keep them beyond the call."""
|
||||
|
||||
def __init__(self, envelope: dict) -> None:
|
||||
ctx = envelope.get("context") or {}
|
||||
cfg = envelope.get("config") or {}
|
||||
self.tenant_id: int = int(ctx.get("tenantId") or 0)
|
||||
self.user_id: str = ctx.get("userId") or ""
|
||||
self.locale: str = ctx.get("locale") or ""
|
||||
self.request_id: str = ctx.get("requestId") or ""
|
||||
#: When WeKnora stops waiting; give up by then.
|
||||
self.deadline: Optional[datetime] = parse_time(ctx.get("deadline"))
|
||||
self._host = ctx.get("host") or {}
|
||||
#: Platform-wide configuration (config.system).
|
||||
self.system: Dict[str, Any] = cfg.get("system") or {}
|
||||
#: Workspace configuration (config.tenant).
|
||||
self.tenant: Dict[str, Any] = cfg.get("tenant") or {}
|
||||
#: The integration instance's configuration, e.g. a data source's.
|
||||
self.instance: Dict[str, Any] = cfg.get("instance") or {}
|
||||
|
||||
def time_left(self) -> Optional[float]:
|
||||
"""Seconds until the deadline, or None without one."""
|
||||
if self.deadline is None:
|
||||
return None
|
||||
return (self.deadline - datetime.now(timezone.utc)).total_seconds()
|
||||
|
||||
def host(self) -> Optional[Host]:
|
||||
"""The Host API client of this call, or None when the plugin was
|
||||
granted no Host API scopes."""
|
||||
if not self._host.get("url") or not self._host.get("token"):
|
||||
return None
|
||||
return Host(self._host["url"], self._host["token"])
|
||||
|
||||
|
||||
class StreamClosed(Exception):
|
||||
"""WeKnora stopped listening to a stream: stop fetching and return."""
|
||||
|
||||
|
||||
class Stream:
|
||||
"""A streaming answer: items, checkpoints and progress as NDJSON lines,
|
||||
each sent at once. Safe to use from several threads."""
|
||||
|
||||
def __init__(self, handler: "_Handler") -> None:
|
||||
self._h = handler
|
||||
self._lock = threading.Lock()
|
||||
self._started = False
|
||||
self._closed = False
|
||||
|
||||
def _write(self, event: dict) -> None:
|
||||
with self._lock:
|
||||
if self._closed:
|
||||
raise StreamClosed("stream already ended")
|
||||
try:
|
||||
if not self._started:
|
||||
self._h.send_response(200)
|
||||
self._h.send_header("Content-Type", p.NDJSON_CONTENT_TYPE)
|
||||
self._h.send_header("Transfer-Encoding", "chunked")
|
||||
self._h.end_headers()
|
||||
self._started = True
|
||||
line = json.dumps(event, separators=(",", ":")).encode() + b"\n"
|
||||
self._h.wfile.write(b"%x\r\n%s\r\n" % (len(line), line))
|
||||
self._h.wfile.flush()
|
||||
except OSError as e:
|
||||
self._closed = True
|
||||
raise StreamClosed(str(e)) from None
|
||||
|
||||
def item(self, item: FetchedItem) -> None:
|
||||
"""Emits one fetched item."""
|
||||
self._write({"type": "item", "data": to_wire(item)})
|
||||
|
||||
def checkpoint(self, cursor: Cursor) -> None:
|
||||
"""Emits a resumable cursor; it must be a complete snapshot."""
|
||||
self._write({"type": "checkpoint", "data": to_wire(cursor)})
|
||||
|
||||
def progress(self, message: str) -> None:
|
||||
"""Reports progress for display."""
|
||||
self._write({"type": "progress", "message": message})
|
||||
|
||||
def log(self, level: str, message: str) -> None:
|
||||
"""Forwards a line to WeKnora's logs: debug, info, warn or error."""
|
||||
self._write({"type": "log", "level": level, "message": message})
|
||||
|
||||
def _finish(self) -> None:
|
||||
try:
|
||||
self._h.wfile.write(b"0\r\n\r\n")
|
||||
self._h.wfile.flush()
|
||||
except OSError:
|
||||
pass
|
||||
self._closed = True
|
||||
|
||||
def _end(self, cursor: Optional[Cursor]) -> None:
|
||||
try:
|
||||
self._write({"type": "end", "data": to_wire(cursor or Cursor())})
|
||||
except StreamClosed:
|
||||
return
|
||||
with self._lock:
|
||||
self._finish()
|
||||
|
||||
def _fail(self, err: PluginError) -> None:
|
||||
with self._lock:
|
||||
if self._closed:
|
||||
return
|
||||
started = self._started
|
||||
if not started:
|
||||
self._h._send_error(err)
|
||||
self._closed = True
|
||||
return
|
||||
try:
|
||||
self._write({"type": "error", "error": err.to_wire()})
|
||||
except StreamClosed:
|
||||
return
|
||||
with self._lock:
|
||||
self._finish()
|
||||
|
||||
|
||||
WebSearchFunc = Callable[[Call, SearchInput], List[SearchResult]]
|
||||
ParserFunc = Callable[[Call, ParseInput], ParseOutput]
|
||||
ConfigValidator = Callable[[Call], None]
|
||||
|
||||
_ROUTE = re.compile(r"^/v1/(websearch|connectors|parsers)/([^/]+)/([a-z-]+)$")
|
||||
|
||||
|
||||
class Plugin:
|
||||
"""A plugin under construction: register contributions, then serve.
|
||||
|
||||
``id`` and ``version`` must match plugin.yaml."""
|
||||
|
||||
def __init__(self, id: str, version: str, logger: Optional[logging.Logger] = None) -> None:
|
||||
self.id = id
|
||||
self.version = version
|
||||
self._web_search: Dict[str, WebSearchFunc] = {}
|
||||
self._connectors: Dict[str, Any] = {}
|
||||
self._parsers: Dict[str, ParserFunc] = {}
|
||||
self._validate: Optional[ConfigValidator] = None
|
||||
#: Seconds serve() waits for calls in flight after SIGTERM.
|
||||
self.shutdown_timeout = 60.0
|
||||
if logger is None:
|
||||
logger = logging.getLogger("weknora_plugin")
|
||||
if not logger.handlers:
|
||||
h = logging.StreamHandler(sys.stderr)
|
||||
h.setFormatter(logging.Formatter("%(asctime)s %(levelname)s %(message)s"))
|
||||
logger.addHandler(h)
|
||||
logger.setLevel(logging.INFO)
|
||||
#: Writes to stderr, which the plugin host forwards to WeKnora's
|
||||
#: logs. Keep stdout for the handshake.
|
||||
self.logger = logger
|
||||
|
||||
# Registration. Each works as a call or as a decorator.
|
||||
|
||||
def web_search(self, id: str, fn: Optional[WebSearchFunc] = None) -> Any:
|
||||
"""Registers web search provider id: fn(call, SearchInput) returns
|
||||
a list of SearchResult."""
|
||||
return self._register(self._web_search, id, fn)
|
||||
|
||||
def parser(self, id: str, fn: Optional[ParserFunc] = None) -> Any:
|
||||
"""Registers parser id: fn(call, ParseInput) returns a ParseOutput
|
||||
(or the Markdown as a str)."""
|
||||
return self._register(self._parsers, id, fn)
|
||||
|
||||
def connector(self, id: str, connector: Any = None) -> Any:
|
||||
"""Registers connector id: an object with validate(call, cfg),
|
||||
list_resources(call, cfg, parent_id) and fetch(call, cfg, input,
|
||||
stream), and optionally resolve_ancestors(call, cfg, resource_ids).
|
||||
As a decorator on a class it registers an instance."""
|
||||
if connector is not None:
|
||||
self._connectors[id] = connector() if isinstance(connector, type) else connector
|
||||
return connector
|
||||
|
||||
def register(c: Any) -> Any:
|
||||
self._connectors[id] = c() if isinstance(c, type) else c
|
||||
return c
|
||||
|
||||
return register
|
||||
|
||||
def config_validator(self, fn: ConfigValidator) -> ConfigValidator:
|
||||
"""Registers the check behind POST /v1/config/validate: raise
|
||||
invalid_config(...) to point at fields."""
|
||||
self._validate = fn
|
||||
return fn
|
||||
|
||||
@staticmethod
|
||||
def _register(table: dict, id: str, fn: Any) -> Any:
|
||||
if fn is not None:
|
||||
table[id] = fn
|
||||
return fn
|
||||
|
||||
def register(f: Any) -> Any:
|
||||
table[id] = f
|
||||
return f
|
||||
|
||||
return register
|
||||
|
||||
def manifest(self) -> dict:
|
||||
"""What the plugin reports at GET /v1/manifest."""
|
||||
contributes = {}
|
||||
for point, table in (
|
||||
("webSearch", self._web_search),
|
||||
("connectors", self._connectors),
|
||||
("parsers", self._parsers),
|
||||
):
|
||||
if table:
|
||||
contributes[point] = sorted(table)
|
||||
return {"id": self.id, "version": self.version, "apiVersion": p.API_VERSION, "contributes": contributes}
|
||||
|
||||
# Dispatch.
|
||||
|
||||
def _dispatch(self, h: "_Handler", method: str, path: str, body: bytes) -> None:
|
||||
if method == "GET" and path == "/v1/manifest":
|
||||
h._send_json(200, self.manifest())
|
||||
return
|
||||
if method == "GET" and path == "/v1/health":
|
||||
h._send_json(200, {"status": "ok"})
|
||||
return
|
||||
if method == "POST" and path == "/v1/config/validate":
|
||||
self._unary(h, body, lambda call, _: self._validate(call) if self._validate else None, empty=True)
|
||||
return
|
||||
m = _ROUTE.match(path)
|
||||
if method == "POST" and m:
|
||||
point, cid, action = m.groups()
|
||||
if self._route(h, point, cid, action, body):
|
||||
return
|
||||
h._send_error(PluginError(ErrorCode.NOT_FOUND, f"no endpoint {method} {path}"))
|
||||
|
||||
def _route(self, h: "_Handler", point: str, cid: str, action: str, body: bytes) -> bool:
|
||||
if point == "websearch" and action == "search":
|
||||
fn = self._web_search.get(cid)
|
||||
if fn is None:
|
||||
h._send_error(PluginError(ErrorCode.NOT_FOUND, f"no web search provider {cid!r}"))
|
||||
else:
|
||||
self._unary(h, body, lambda call, raw: {"results": list(fn(call, from_wire(SearchInput, raw)) or [])})
|
||||
return True
|
||||
if point == "parsers" and action == "parse":
|
||||
fn = self._parsers.get(cid)
|
||||
if fn is None:
|
||||
h._send_error(PluginError(ErrorCode.NOT_FOUND, f"no parser {cid!r}"))
|
||||
else:
|
||||
self._unary(h, body, lambda call, raw: _parse_output(fn(call, from_wire(ParseInput, raw))))
|
||||
return True
|
||||
if point == "connectors" and action in ("validate", "list-resources", "resolve-ancestors", "fetch"):
|
||||
c = self._connectors.get(cid)
|
||||
if c is None:
|
||||
h._send_error(PluginError(ErrorCode.NOT_FOUND, f"no connector {cid!r}"))
|
||||
elif action == "fetch":
|
||||
self._fetch(h, c, body)
|
||||
else:
|
||||
self._unary(h, body, lambda call, raw: _connector_call(c, action, call, raw), empty=action == "validate")
|
||||
return True
|
||||
return False
|
||||
|
||||
def _unary(self, h: "_Handler", body: bytes, fn: Callable[[Call, Any], Any], empty: bool = False) -> None:
|
||||
try:
|
||||
env = _envelope(body)
|
||||
out = fn(Call(env), env.get("input"))
|
||||
h._send_json(200, {"output": {} if empty else to_wire(out)})
|
||||
except Exception as e: # noqa: BLE001 - every failure becomes a protocol error
|
||||
h._send_error(self._as_error(e))
|
||||
|
||||
def _fetch(self, h: "_Handler", c: Any, body: bytes) -> None:
|
||||
try:
|
||||
env = _envelope(body)
|
||||
call = Call(env)
|
||||
cfg = _connector_config(call)
|
||||
inp = from_wire(FetchInput, env.get("input"))
|
||||
except Exception as e: # noqa: BLE001
|
||||
h._send_error(self._as_error(e))
|
||||
return
|
||||
s = Stream(h)
|
||||
try:
|
||||
cursor = c.fetch(call, cfg, inp, s)
|
||||
except StreamClosed:
|
||||
return
|
||||
except Exception as e: # noqa: BLE001
|
||||
s._fail(self._as_error(e))
|
||||
return
|
||||
s._end(cursor)
|
||||
|
||||
def _as_error(self, e: Exception) -> PluginError:
|
||||
if isinstance(e, PluginError):
|
||||
return e
|
||||
self.logger.error("plugin call failed: %s", "".join(traceback.format_exception(type(e), e, e.__traceback__)))
|
||||
return PluginError(ErrorCode.INTERNAL, str(e) or type(e).__name__)
|
||||
|
||||
# Serving.
|
||||
|
||||
def serve(self) -> None:
|
||||
"""Runs the plugin until SIGTERM or SIGINT, then drains calls in
|
||||
flight. Started by the plugin host it listens where the host says
|
||||
and prints the handshake; otherwise it serves as a remote plugin on
|
||||
WEKNORA_PLUGIN_ADDR, checking signatures with WEKNORA_PLUGIN_SECRET."""
|
||||
server, handshake = self._listen()
|
||||
stop = threading.Event()
|
||||
|
||||
def on_signal(*_: Any) -> None:
|
||||
stop.set()
|
||||
|
||||
for sig in (signal.SIGTERM, signal.SIGINT):
|
||||
try:
|
||||
signal.signal(sig, on_signal)
|
||||
except ValueError: # not the main thread
|
||||
pass
|
||||
t = threading.Thread(target=server.serve_forever, kwargs={"poll_interval": 0.2}, daemon=True)
|
||||
t.start()
|
||||
if handshake:
|
||||
# The host waits for this exact line; everything else on stdout
|
||||
# is treated as log output.
|
||||
print(handshake, flush=True)
|
||||
else:
|
||||
self.logger.info("plugin serving on %s", _describe(server))
|
||||
stop.wait()
|
||||
server.shutdown()
|
||||
server.drain(self.shutdown_timeout)
|
||||
server.server_close()
|
||||
|
||||
def _listen(self) -> Tuple["_Server", Optional[str]]:
|
||||
sock_path = os.environ.get(p.ENV_SOCKET, "")
|
||||
if sock_path:
|
||||
network = os.environ.get(p.ENV_NETWORK) or "unix"
|
||||
token = os.environ.get(p.ENV_TOKEN, "")
|
||||
if not token:
|
||||
raise RuntimeError(f"{p.ENV_SOCKET} is set without {p.ENV_TOKEN}")
|
||||
auth = _bearer(token)
|
||||
if network == "unix":
|
||||
try:
|
||||
os.remove(sock_path)
|
||||
except FileNotFoundError:
|
||||
pass
|
||||
server: _Server = _UnixServer(sock_path, self, auth)
|
||||
return server, p.handshake_line("unix", sock_path)
|
||||
server = _TCPServer(_split_addr(sock_path), self, auth)
|
||||
host, port = server.server_address[:2]
|
||||
return server, p.handshake_line("tcp", f"{host}:{port}")
|
||||
secret = os.environ.get(p.ENV_SECRET, "")
|
||||
if not secret:
|
||||
raise RuntimeError(f"a remote plugin needs {p.ENV_SECRET} (the secret shown when it was registered)")
|
||||
addr = os.environ.get(p.ENV_ADDR) or ":8080"
|
||||
return _TCPServer(_split_addr(addr), self, _signed(secret.encode())), None
|
||||
|
||||
def test_server(self, auth: Optional[Callable[["_Handler", bytes], Optional[str]]] = None) -> "_Server":
|
||||
"""Serves the plugin on a free loopback port without authentication
|
||||
(or with the given check), for tests. Call shutdown() when done."""
|
||||
server = _TCPServer(("127.0.0.1", 0), self, auth or (lambda h, b: None))
|
||||
threading.Thread(target=server.serve_forever, kwargs={"poll_interval": 0.1}, daemon=True).start()
|
||||
return server
|
||||
|
||||
|
||||
def _envelope(body: bytes) -> dict:
|
||||
try:
|
||||
env = json.loads(body or b"{}")
|
||||
except ValueError as e:
|
||||
raise PluginError(ErrorCode.BAD_REQUEST, f"decode envelope: {e}") from None
|
||||
if not isinstance(env, dict):
|
||||
raise PluginError(ErrorCode.BAD_REQUEST, "decode envelope: not an object")
|
||||
return env
|
||||
|
||||
|
||||
def _connector_config(call: Call) -> ConnectorConfig:
|
||||
try:
|
||||
return from_wire(ConnectorConfig, call.instance)
|
||||
except (TypeError, ValueError) as e:
|
||||
raise PluginError(ErrorCode.BAD_REQUEST, f"decode connector config: {e}") from None
|
||||
|
||||
|
||||
def _connector_call(c: Any, action: str, call: Call, raw: Any) -> Any:
|
||||
cfg = _connector_config(call)
|
||||
raw = raw or {}
|
||||
if action == "validate":
|
||||
c.validate(call, cfg)
|
||||
return None
|
||||
if action == "list-resources":
|
||||
return {"resources": list(c.list_resources(call, cfg, raw.get("parentId") or "") or [])}
|
||||
resolve = getattr(c, "resolve_ancestors", None)
|
||||
if resolve is None:
|
||||
return {"ancestors": []}
|
||||
return {"ancestors": list(resolve(call, cfg, list(raw.get("resourceIds") or [])) or [])}
|
||||
|
||||
|
||||
def _parse_output(out: Any) -> Any:
|
||||
if isinstance(out, str):
|
||||
return ParseOutput(markdown=out)
|
||||
return out
|
||||
|
||||
|
||||
# Authentication: a check returns an error message, or None to let the
|
||||
# request through.
|
||||
|
||||
|
||||
def _bearer(token: str) -> Callable[["_Handler", bytes], Optional[str]]:
|
||||
want = "Bearer " + token
|
||||
|
||||
def check(h: "_Handler", _: bytes) -> Optional[str]:
|
||||
if hmac.compare_digest(h.headers.get("Authorization", ""), want):
|
||||
return None
|
||||
return "missing or wrong host token"
|
||||
|
||||
return check
|
||||
|
||||
|
||||
def _signed(secret: bytes) -> Callable[["_Handler", bytes], Optional[str]]:
|
||||
def check(h: "_Handler", body: bytes) -> Optional[str]:
|
||||
try:
|
||||
p.verify_signature(
|
||||
secret, h.headers.get(p.TIMESTAMP_HEADER, ""), h.headers.get(p.SIGNATURE_HEADER, ""), body
|
||||
)
|
||||
except ValueError as e:
|
||||
return str(e)
|
||||
return None
|
||||
|
||||
return check
|
||||
|
||||
|
||||
class _Handler(BaseHTTPRequestHandler):
|
||||
protocol_version = "HTTP/1.1"
|
||||
server: "_Server"
|
||||
|
||||
def _handle(self) -> None:
|
||||
length = int(self.headers.get("Content-Length") or 0)
|
||||
body = self.rfile.read(length) if length > 0 else b""
|
||||
path = self.path.split("?", 1)[0]
|
||||
problem = self.server.auth(self, body)
|
||||
if problem is not None:
|
||||
self._send_error(PluginError(ErrorCode.UNAUTHORIZED, problem))
|
||||
return
|
||||
with self.server.in_flight():
|
||||
self.server.plugin._dispatch(self, self.command, path, body)
|
||||
|
||||
do_GET = do_POST = do_PUT = do_DELETE = do_PATCH = _handle
|
||||
|
||||
def _send_json(self, status: int, value: Any) -> None:
|
||||
raw = json.dumps(value, separators=(",", ":")).encode()
|
||||
self.send_response(status)
|
||||
self.send_header("Content-Type", "application/json")
|
||||
self.send_header("Content-Length", str(len(raw)))
|
||||
self.end_headers()
|
||||
self.wfile.write(raw)
|
||||
|
||||
def _send_error(self, err: PluginError) -> None:
|
||||
self._send_json(err.code.http_status, {"error": err.to_wire()})
|
||||
|
||||
def address_string(self) -> str:
|
||||
return str(self.client_address[0]) if self.client_address else "host"
|
||||
|
||||
def log_message(self, format: str, *args: Any) -> None:
|
||||
pass
|
||||
|
||||
|
||||
class _Server(socketserver.ThreadingMixIn, HTTPServer):
|
||||
daemon_threads = True
|
||||
|
||||
def __init__(self, address: Any, plugin: Plugin, auth: Callable[[_Handler, bytes], Optional[str]]) -> None:
|
||||
self.plugin = plugin
|
||||
self.auth = auth
|
||||
self._active = 0
|
||||
self._active_lock = threading.Condition()
|
||||
super().__init__(address, _Handler)
|
||||
|
||||
@contextlib.contextmanager
|
||||
def in_flight(self) -> Iterator[None]:
|
||||
"""Counts a call being answered; idle keep-alive connections do not
|
||||
hold up a shutdown."""
|
||||
with self._active_lock:
|
||||
self._active += 1
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
with self._active_lock:
|
||||
self._active -= 1
|
||||
self._active_lock.notify_all()
|
||||
|
||||
def handle_error(self, request: Any, client_address: Any) -> None:
|
||||
# A caller that went away mid-answer is routine.
|
||||
self.plugin.logger.debug("connection error", exc_info=True)
|
||||
|
||||
def drain(self, timeout: float) -> None:
|
||||
"""Waits for calls in flight, up to timeout seconds."""
|
||||
end = time.monotonic() + timeout
|
||||
with self._active_lock:
|
||||
while self._active and time.monotonic() < end:
|
||||
self._active_lock.wait(end - time.monotonic())
|
||||
|
||||
|
||||
class _TCPServer(_Server):
|
||||
address_family = socket.AF_INET
|
||||
|
||||
def __init__(self, address: Tuple[str, int], *args: Any) -> None:
|
||||
if ":" in address[0]:
|
||||
self.address_family = socket.AF_INET6
|
||||
super().__init__(address, *args)
|
||||
|
||||
def server_bind(self) -> None:
|
||||
# HTTPServer.server_bind looks the host up with getfqdn, a reverse
|
||||
# DNS query that can stall startup for seconds; nothing needs it.
|
||||
socketserver.TCPServer.server_bind(self)
|
||||
self.server_name, self.server_port = str(self.server_address[0]), int(self.server_address[1])
|
||||
|
||||
|
||||
class _UnixServer(_Server):
|
||||
address_family = getattr(socket, "AF_UNIX", socket.AF_INET)
|
||||
|
||||
def server_bind(self) -> None:
|
||||
socketserver.TCPServer.server_bind(self)
|
||||
self.server_name, self.server_port = "plugin", 0
|
||||
|
||||
def server_close(self) -> None:
|
||||
super().server_close()
|
||||
try:
|
||||
os.remove(self.server_address)
|
||||
except (OSError, TypeError):
|
||||
pass
|
||||
|
||||
|
||||
def _split_addr(addr: str) -> Tuple[str, int]:
|
||||
host, _, port = addr.rpartition(":")
|
||||
host = host.strip("[]")
|
||||
return (host or "0.0.0.0", int(port))
|
||||
|
||||
|
||||
def _describe(server: _Server) -> str:
|
||||
addr = server.server_address
|
||||
return f"{addr[0]}:{addr[1]}" if isinstance(addr, tuple) else str(addr)
|
||||
@@ -0,0 +1,149 @@
|
||||
"""Version 1 of the WeKnora extension protocol, as the Go package pluginapi
|
||||
defines it: constants, errors, request signatures and the handshake line."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import hmac
|
||||
import time
|
||||
from enum import Enum
|
||||
from typing import Dict, Optional
|
||||
|
||||
PROTOCOL_VERSION = "1"
|
||||
API_VERSION = "weknora.plugin/v1"
|
||||
|
||||
PROTOCOL_HEADER = "X-WeKnora-Protocol"
|
||||
REQUEST_ID_HEADER = "X-Request-Id"
|
||||
SIGNATURE_HEADER = "X-WeKnora-Signature"
|
||||
TIMESTAMP_HEADER = "X-WeKnora-Timestamp"
|
||||
NDJSON_CONTENT_TYPE = "application/x-ndjson"
|
||||
|
||||
# Environment a host plugin is started with.
|
||||
ENV_SOCKET = "WEKNORA_PLUGIN_SOCKET"
|
||||
ENV_NETWORK = "WEKNORA_PLUGIN_NETWORK"
|
||||
ENV_TOKEN = "WEKNORA_PLUGIN_TOKEN"
|
||||
ENV_PLUGIN_ID = "WEKNORA_PLUGIN_ID"
|
||||
ENV_PLUGIN_VERSION = "WEKNORA_PLUGIN_VERSION"
|
||||
# Environment of a remote plugin.
|
||||
ENV_ADDR = "WEKNORA_PLUGIN_ADDR"
|
||||
ENV_SECRET = "WEKNORA_PLUGIN_SECRET"
|
||||
|
||||
# How old a signed request may be, in seconds.
|
||||
MAX_CLOCK_SKEW = 5 * 60
|
||||
|
||||
HOST_KV_PATH = "/api/v1/plugin-host/kv"
|
||||
HOST_KV_LIST_PATH = "/api/v1/plugin-host/kv/list"
|
||||
|
||||
|
||||
class ErrorCode(str, Enum):
|
||||
"""Classifies a failed call; WeKnora reacts to each differently."""
|
||||
|
||||
INVALID_CONFIG = "invalid_config"
|
||||
UNAUTHORIZED = "unauthorized"
|
||||
NOT_FOUND = "not_found"
|
||||
RATE_LIMITED = "rate_limited"
|
||||
UNAVAILABLE = "unavailable"
|
||||
BAD_REQUEST = "bad_request"
|
||||
INTERNAL = "internal"
|
||||
|
||||
@property
|
||||
def http_status(self) -> int:
|
||||
return _STATUS.get(self, 500)
|
||||
|
||||
|
||||
_STATUS = {
|
||||
ErrorCode.INVALID_CONFIG: 400,
|
||||
ErrorCode.BAD_REQUEST: 400,
|
||||
ErrorCode.UNAUTHORIZED: 401,
|
||||
ErrorCode.NOT_FOUND: 404,
|
||||
ErrorCode.RATE_LIMITED: 429,
|
||||
ErrorCode.UNAVAILABLE: 503,
|
||||
}
|
||||
|
||||
|
||||
class PluginError(Exception):
|
||||
"""A failed call. Raise it to choose the code WeKnora sees; any other
|
||||
exception becomes ``internal``. Unavailable and rate-limited errors are
|
||||
retryable unless told otherwise."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
code: ErrorCode | str,
|
||||
message: str,
|
||||
*,
|
||||
retryable: Optional[bool] = None,
|
||||
fields: Optional[Dict[str, str]] = None,
|
||||
retry_after: int = 0,
|
||||
) -> None:
|
||||
super().__init__(message)
|
||||
try:
|
||||
self.code = ErrorCode(code)
|
||||
except ValueError:
|
||||
self.code = ErrorCode.INTERNAL
|
||||
self.message = message
|
||||
if retryable is None:
|
||||
retryable = self.code in (ErrorCode.UNAVAILABLE, ErrorCode.RATE_LIMITED)
|
||||
self.retryable = retryable
|
||||
self.fields = fields or {}
|
||||
self.retry_after = retry_after
|
||||
|
||||
def __str__(self) -> str:
|
||||
return f"{self.code.value}: {self.message}"
|
||||
|
||||
def to_wire(self) -> dict:
|
||||
out: dict = {"code": self.code.value, "message": self.message}
|
||||
if self.retryable:
|
||||
out["retryable"] = True
|
||||
details: dict = {}
|
||||
if self.fields:
|
||||
details["fields"] = dict(self.fields)
|
||||
if self.retry_after:
|
||||
details["retryAfter"] = self.retry_after
|
||||
if details:
|
||||
out["details"] = details
|
||||
return out
|
||||
|
||||
@classmethod
|
||||
def from_wire(cls, body: dict) -> "PluginError":
|
||||
details = body.get("details") or {}
|
||||
return cls(
|
||||
body.get("code") or ErrorCode.INTERNAL,
|
||||
body.get("message") or "",
|
||||
retryable=bool(body.get("retryable")),
|
||||
fields=details.get("fields"),
|
||||
retry_after=int(details.get("retryAfter") or 0),
|
||||
)
|
||||
|
||||
|
||||
def invalid_config(message: str, fields: Dict[str, str]) -> PluginError:
|
||||
"""Reports configuration problems keyed by dotted field path; WeKnora
|
||||
shows them on the form."""
|
||||
return PluginError(ErrorCode.INVALID_CONFIG, message, fields=fields)
|
||||
|
||||
|
||||
def sign(secret: bytes, timestamp: int, body: bytes) -> str:
|
||||
"""The signature of a request body at a unix timestamp: hex HMAC-SHA256
|
||||
over ``"<timestamp>.<body>"``."""
|
||||
mac = hmac.new(secret, f"{timestamp}.".encode(), hashlib.sha256)
|
||||
mac.update(body)
|
||||
return mac.hexdigest()
|
||||
|
||||
|
||||
def verify_signature(
|
||||
secret: bytes, timestamp: str, signature: str, body: bytes, now: Optional[float] = None
|
||||
) -> None:
|
||||
"""Checks a signed request; raises ValueError when it does not hold."""
|
||||
try:
|
||||
ts = int(timestamp)
|
||||
except (TypeError, ValueError):
|
||||
raise ValueError(f"bad {TIMESTAMP_HEADER}") from None
|
||||
now = time.time() if now is None else now
|
||||
if abs(now - ts) > MAX_CLOCK_SKEW:
|
||||
raise ValueError("request timestamp is outside the allowed clock skew")
|
||||
if not hmac.compare_digest(sign(secret, ts, body), signature or ""):
|
||||
raise ValueError("bad signature")
|
||||
|
||||
|
||||
def handshake_line(network: str, address: str) -> str:
|
||||
"""The line a host plugin prints on stdout once it is ready."""
|
||||
return "|".join(["WEKNORA_PLUGIN", PROTOCOL_VERSION, network, address])
|
||||
@@ -0,0 +1,242 @@
|
||||
"""The protocol's messages as dataclasses, and their JSON form: camelCase
|
||||
keys, bytes as base64, times as RFC 3339. Optional fields left at their
|
||||
default are omitted, like Go's omitempty; unknown fields are ignored."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import dataclasses
|
||||
import re
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any, Dict, List, Optional, Type, TypeVar, Union
|
||||
|
||||
T = TypeVar("T")
|
||||
|
||||
|
||||
def _camel(name: str) -> str:
|
||||
head, *rest = name.split("_")
|
||||
return head + "".join(p[:1].upper() + p[1:] for p in rest)
|
||||
|
||||
|
||||
_FRACTION = re.compile(r"(\.\d{6})\d+")
|
||||
|
||||
|
||||
def parse_time(value: Any) -> Optional[datetime]:
|
||||
"""Reads an RFC 3339 time as Go writes it (nanoseconds, a Z suffix)."""
|
||||
if not value:
|
||||
return None
|
||||
if isinstance(value, datetime):
|
||||
return value
|
||||
s = _FRACTION.sub(r"\1", str(value).replace("Z", "+00:00"))
|
||||
return datetime.fromisoformat(s)
|
||||
|
||||
|
||||
def format_time(value: datetime) -> str:
|
||||
if value.tzinfo is None:
|
||||
value = value.replace(tzinfo=timezone.utc)
|
||||
return value.astimezone(timezone.utc).isoformat().replace("+00:00", "Z")
|
||||
|
||||
|
||||
def to_wire(value: Any) -> Any:
|
||||
"""Converts a message (or anything JSON-like holding messages) to its
|
||||
JSON form."""
|
||||
if dataclasses.is_dataclass(value) and not isinstance(value, type):
|
||||
out = {}
|
||||
for f in dataclasses.fields(value):
|
||||
v = getattr(value, f.name)
|
||||
optional = f.default is not dataclasses.MISSING or f.default_factory is not dataclasses.MISSING
|
||||
if optional and (v is None or v is False or v == "" or v == b"" or v == [] or v == {}):
|
||||
continue
|
||||
out[_camel(f.name)] = to_wire(v)
|
||||
return out
|
||||
if isinstance(value, (bytes, bytearray)):
|
||||
return base64.b64encode(value).decode()
|
||||
if isinstance(value, datetime):
|
||||
return format_time(value)
|
||||
if isinstance(value, dict):
|
||||
return {k: to_wire(v) for k, v in value.items()}
|
||||
if isinstance(value, (list, tuple)):
|
||||
return [to_wire(v) for v in value]
|
||||
return value
|
||||
|
||||
|
||||
def from_wire(cls: Type[T], data: Any) -> T:
|
||||
"""Builds a message from its JSON form."""
|
||||
if isinstance(data, cls):
|
||||
return data
|
||||
data = data or {}
|
||||
kwargs = {}
|
||||
hints = _hints(cls)
|
||||
for f in dataclasses.fields(cls):
|
||||
key = _camel(f.name)
|
||||
if key not in data:
|
||||
continue
|
||||
kwargs[f.name] = _convert(hints[f.name], data[key])
|
||||
return cls(**kwargs)
|
||||
|
||||
|
||||
def _hints(cls: type) -> Dict[str, Any]:
|
||||
import typing
|
||||
|
||||
return typing.get_type_hints(cls)
|
||||
|
||||
|
||||
def _convert(hint: Any, value: Any) -> Any:
|
||||
import typing
|
||||
|
||||
if value is None:
|
||||
return None
|
||||
origin = typing.get_origin(hint)
|
||||
args = typing.get_args(hint)
|
||||
if origin is Union:
|
||||
inner = [a for a in args if a is not type(None)]
|
||||
if len(inner) == 1:
|
||||
return _convert(inner[0], value)
|
||||
if bytes in inner and isinstance(value, str):
|
||||
return base64.b64decode(value)
|
||||
return value
|
||||
if hint is bytes:
|
||||
return base64.b64decode(value) if isinstance(value, str) else bytes(value)
|
||||
if hint is datetime:
|
||||
return parse_time(value)
|
||||
if dataclasses.is_dataclass(hint):
|
||||
return from_wire(hint, value)
|
||||
if origin in (list, List) and args and isinstance(value, list):
|
||||
return [_convert(args[0], v) for v in value]
|
||||
return value
|
||||
|
||||
|
||||
@dataclass
|
||||
class ConnectorConfig:
|
||||
"""A data source's instance configuration. Credentials are decrypted;
|
||||
do not keep them beyond the call."""
|
||||
|
||||
credentials: Dict[str, Any] = field(default_factory=dict)
|
||||
settings: Dict[str, Any] = field(default_factory=dict)
|
||||
resource_ids: List[str] = field(default_factory=list)
|
||||
|
||||
|
||||
@dataclass
|
||||
class Resource:
|
||||
"""Something a data source can sync: a space, a folder, a feed."""
|
||||
|
||||
external_id: str
|
||||
name: str
|
||||
type: str = ""
|
||||
description: str = ""
|
||||
url: str = ""
|
||||
modified_at: Optional[datetime] = None
|
||||
parent_id: str = ""
|
||||
has_children: bool = False
|
||||
metadata: Optional[Dict[str, Any]] = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class Cursor:
|
||||
"""A connector's resumable state; WeKnora stores it and hands it back.
|
||||
A checkpoint must be a complete snapshot of progress so far."""
|
||||
|
||||
last_sync_time: Optional[datetime] = None
|
||||
state: Optional[Dict[str, Any]] = None
|
||||
|
||||
|
||||
FETCH_FULL = "full"
|
||||
FETCH_INCREMENTAL = "incremental"
|
||||
|
||||
|
||||
@dataclass
|
||||
class FetchInput:
|
||||
"""Starts or resumes a sync; cursor None starts from scratch."""
|
||||
|
||||
mode: str = FETCH_FULL
|
||||
cursor: Optional[Cursor] = None
|
||||
resource_ids: List[str] = field(default_factory=list)
|
||||
|
||||
|
||||
@dataclass
|
||||
class FetchedItem:
|
||||
"""One document fetched from the source. Content is the body, Markdown
|
||||
preferred; a str is sent as UTF-8."""
|
||||
|
||||
external_id: str
|
||||
title: str
|
||||
content: Optional[Union[bytes, str]] = None
|
||||
content_type: str = ""
|
||||
file_name: str = ""
|
||||
url: str = ""
|
||||
updated_at: Optional[datetime] = None
|
||||
created_at: Optional[datetime] = None
|
||||
metadata: Optional[Dict[str, str]] = None
|
||||
is_deleted: bool = False
|
||||
source_resource_id: str = ""
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
if isinstance(self.content, str):
|
||||
self.content = self.content.encode()
|
||||
|
||||
|
||||
@dataclass
|
||||
class SearchInput:
|
||||
"""A web search. A provider that cannot honour region or freshness must
|
||||
fail rather than ignore it."""
|
||||
|
||||
query: str = ""
|
||||
max_results: int = 0
|
||||
include_date: bool = False
|
||||
region: str = ""
|
||||
freshness: str = ""
|
||||
|
||||
|
||||
@dataclass
|
||||
class SearchResult:
|
||||
title: str
|
||||
url: str
|
||||
snippet: str = ""
|
||||
content: str = ""
|
||||
age: str = ""
|
||||
published_at: Optional[datetime] = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class ParseInput:
|
||||
"""One document to parse: its bytes, or a URL. file_type is the
|
||||
lower-case extension without the dot."""
|
||||
|
||||
file_name: str = ""
|
||||
file_type: str = ""
|
||||
content: bytes = b""
|
||||
url: str = ""
|
||||
title: str = ""
|
||||
|
||||
|
||||
@dataclass
|
||||
class ParsedImage:
|
||||
"""An image the Markdown references as ."""
|
||||
|
||||
original_ref: str
|
||||
data: bytes
|
||||
mime_type: str = ""
|
||||
|
||||
|
||||
@dataclass
|
||||
class ParseOutput:
|
||||
"""A document as Markdown; WeKnora chunks it."""
|
||||
|
||||
markdown: str
|
||||
images: List[ParsedImage] = field(default_factory=list)
|
||||
metadata: Optional[Dict[str, str]] = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class KVEntry:
|
||||
key: str
|
||||
value: Any = None
|
||||
expires_at: Optional[datetime] = None
|
||||
updated_at: Optional[datetime] = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class KVList:
|
||||
entries: List[KVEntry] = field(default_factory=list)
|
||||
next: str = ""
|
||||
@@ -0,0 +1,78 @@
|
||||
"""A plugin exercising every contribution kind: the SDK's tests and the Go
|
||||
conformance suite run it."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import sys
|
||||
import time
|
||||
|
||||
sys.path.insert(0, os.path.join(os.path.dirname(os.path.abspath(__file__)), "..", "src"))
|
||||
|
||||
from weknora_plugin import ( # noqa: E402
|
||||
Cursor,
|
||||
ErrorCode,
|
||||
FetchedItem,
|
||||
ParsedImage,
|
||||
ParseOutput,
|
||||
Plugin,
|
||||
PluginError,
|
||||
Resource,
|
||||
SearchResult,
|
||||
invalid_config,
|
||||
)
|
||||
|
||||
plugin = Plugin("acme.fixture", "1.0.0")
|
||||
|
||||
|
||||
@plugin.config_validator
|
||||
def validate(call):
|
||||
if call.tenant.get("region") == "mars":
|
||||
raise invalid_config("unknown region", {"region": "no such region"})
|
||||
|
||||
|
||||
@plugin.web_search("echo")
|
||||
def search(call, q):
|
||||
if q.query == "down":
|
||||
raise PluginError(ErrorCode.UNAVAILABLE, "upstream is down")
|
||||
if q.query == "crash":
|
||||
raise RuntimeError("boom")
|
||||
if q.query == "slow":
|
||||
time.sleep(0.5)
|
||||
return [SearchResult(title=q.query, url=f"https://example.com/?tenant={call.tenant_id}", snippet=call.locale)]
|
||||
|
||||
|
||||
@plugin.parser("upper")
|
||||
def parse(call, doc):
|
||||
text = doc.content.decode("utf-8", "replace")
|
||||
return ParseOutput(
|
||||
markdown=text.upper() + "\n\n",
|
||||
images=[ParsedImage(original_ref="img/dot.png", data=b"\x89PNG", mime_type="image/png")],
|
||||
metadata={"fileType": doc.file_type},
|
||||
)
|
||||
|
||||
|
||||
@plugin.connector("notes")
|
||||
class Notes:
|
||||
def validate(self, call, cfg):
|
||||
if cfg.credentials.get("token") == "bad":
|
||||
raise PluginError(ErrorCode.UNAUTHORIZED, "token rejected")
|
||||
|
||||
def list_resources(self, call, cfg, parent_id):
|
||||
return [Resource(external_id="inbox", name="Inbox", has_children=not parent_id)]
|
||||
|
||||
def fetch(self, call, cfg, inp, stream):
|
||||
start = int((inp.cursor.state or {}).get("after", 0)) if inp.cursor else 0
|
||||
if cfg.settings.get("fail") == "before":
|
||||
raise PluginError(ErrorCode.UNAUTHORIZED, "token rejected")
|
||||
for i in range(start, start + 3):
|
||||
stream.item(FetchedItem(external_id=f"n{i}", title=f"Note {i}", content=f"# Note {i}"))
|
||||
if cfg.settings.get("fail") == "during":
|
||||
raise PluginError(ErrorCode.RATE_LIMITED, "slow down", retry_after=7)
|
||||
stream.checkpoint(Cursor(state={"after": start + 3}))
|
||||
stream.progress("done")
|
||||
return Cursor(state={"after": start + 3})
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
plugin.serve()
|
||||
@@ -0,0 +1,363 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import http.client
|
||||
import json
|
||||
import os
|
||||
import socket
|
||||
import subprocess
|
||||
import sys
|
||||
import threading
|
||||
import time
|
||||
import unittest
|
||||
from datetime import datetime, timezone
|
||||
from http.server import BaseHTTPRequestHandler, HTTPServer
|
||||
from urllib.parse import parse_qs, urlparse
|
||||
|
||||
HERE = os.path.dirname(os.path.abspath(__file__))
|
||||
sys.path.insert(0, HERE)
|
||||
|
||||
from fixture import plugin # noqa: E402
|
||||
from weknora_plugin import ErrorCode, Host, PluginError # noqa: E402
|
||||
from weknora_plugin import protocol as p # noqa: E402
|
||||
from weknora_plugin.types import Cursor, FetchInput, from_wire, parse_time, to_wire # noqa: E402
|
||||
|
||||
ENVELOPE = {"context": {"tenantId": 7, "locale": "zh-CN", "requestId": "r1"}}
|
||||
|
||||
|
||||
class Client:
|
||||
def __init__(self, conn_factory):
|
||||
self.conn_factory = conn_factory
|
||||
|
||||
def request(self, method, path, body=None, headers=None):
|
||||
conn = self.conn_factory()
|
||||
raw = b"" if body is None else (body if isinstance(body, bytes) else json.dumps(body).encode())
|
||||
conn.request(method, path, body=raw if method != "GET" else None, headers=headers or {})
|
||||
resp = conn.getresponse()
|
||||
data = resp.read()
|
||||
conn.close()
|
||||
return resp.status, resp.getheader("Content-Type"), data
|
||||
|
||||
def call(self, path, input=None, config=None, **kw):
|
||||
env = dict(ENVELOPE)
|
||||
if input is not None:
|
||||
env["input"] = input
|
||||
if config is not None:
|
||||
env["config"] = config
|
||||
status, _, data = self.request("POST", path, env, **kw)
|
||||
return status, json.loads(data)
|
||||
|
||||
def stream(self, path, input=None, config=None):
|
||||
env = dict(ENVELOPE, input=input or {}, config=config or {})
|
||||
status, ctype, data = self.request("POST", path, env)
|
||||
if ctype != p.NDJSON_CONTENT_TYPE:
|
||||
return status, json.loads(data)
|
||||
return status, [json.loads(line) for line in data.splitlines() if line]
|
||||
|
||||
|
||||
class PluginTest(unittest.TestCase):
|
||||
@classmethod
|
||||
def setUpClass(cls):
|
||||
cls.server = plugin.test_server()
|
||||
host, port = cls.server.server_address[:2]
|
||||
cls.c = Client(lambda: http.client.HTTPConnection(host, port, timeout=10))
|
||||
|
||||
@classmethod
|
||||
def tearDownClass(cls):
|
||||
cls.server.shutdown()
|
||||
cls.server.server_close()
|
||||
|
||||
def test_manifest_and_health(self):
|
||||
status, _, data = self.c.request("GET", "/v1/manifest")
|
||||
self.assertEqual(status, 200)
|
||||
self.assertEqual(
|
||||
json.loads(data),
|
||||
{
|
||||
"id": "acme.fixture",
|
||||
"version": "1.0.0",
|
||||
"apiVersion": "weknora.plugin/v1",
|
||||
"contributes": {"webSearch": ["echo"], "connectors": ["notes"], "parsers": ["upper"]},
|
||||
},
|
||||
)
|
||||
self.assertEqual(json.loads(self.c.request("GET", "/v1/health")[2]), {"status": "ok"})
|
||||
|
||||
def test_web_search(self):
|
||||
status, body = self.c.call("/v1/websearch/echo/search", {"query": "hi", "maxResults": 3})
|
||||
self.assertEqual(status, 200)
|
||||
self.assertEqual(
|
||||
body["output"]["results"],
|
||||
[{"title": "hi", "url": "https://example.com/?tenant=7", "snippet": "zh-CN"}],
|
||||
)
|
||||
|
||||
def test_errors(self):
|
||||
status, body = self.c.call("/v1/websearch/echo/search", {"query": "down"})
|
||||
self.assertEqual((status, body["error"]), (503, {"code": "unavailable", "message": "upstream is down", "retryable": True}))
|
||||
status, body = self.c.call("/v1/websearch/echo/search", {"query": "crash"})
|
||||
self.assertEqual((status, body["error"]["code"], body["error"]["message"]), (500, "internal", "boom"))
|
||||
status, body = self.c.call("/v1/websearch/nope/search", {"query": "x"})
|
||||
self.assertEqual((status, body["error"]["code"]), (404, "not_found"))
|
||||
status, body = self.c.call("/v1/no-such-endpoint")
|
||||
self.assertEqual((status, body["error"]["code"]), (404, "not_found"))
|
||||
status, _, data = self.c.request("POST", "/v1/config/validate", b"{not json")
|
||||
self.assertEqual((status, json.loads(data)["error"]["code"]), (400, "bad_request"))
|
||||
|
||||
def test_config_validate(self):
|
||||
status, body = self.c.call("/v1/config/validate", config={"tenant": {"region": "eu"}})
|
||||
self.assertEqual((status, body), (200, {"output": {}}))
|
||||
status, body = self.c.call("/v1/config/validate", config={"tenant": {"region": "mars"}})
|
||||
self.assertEqual(status, 400)
|
||||
self.assertEqual(body["error"]["details"], {"fields": {"region": "no such region"}})
|
||||
|
||||
def test_parser(self):
|
||||
doc = base64.b64encode(b"hello").decode()
|
||||
status, body = self.c.call("/v1/parsers/upper/parse", {"fileName": "a.txt", "fileType": "txt", "content": doc})
|
||||
self.assertEqual(status, 200)
|
||||
out = body["output"]
|
||||
self.assertTrue(out["markdown"].startswith("HELLO"))
|
||||
self.assertEqual(out["images"], [{"originalRef": "img/dot.png", "data": base64.b64encode(b"\x89PNG").decode(), "mimeType": "image/png"}])
|
||||
self.assertEqual(out["metadata"], {"fileType": "txt"})
|
||||
|
||||
def test_connector_unary(self):
|
||||
cfg = {"instance": {"credentials": {"token": "ok"}, "settings": {}, "resourceIds": []}}
|
||||
self.assertEqual(self.c.call("/v1/connectors/notes/validate", config=cfg), (200, {"output": {}}))
|
||||
bad = {"instance": {"credentials": {"token": "bad"}}}
|
||||
status, body = self.c.call("/v1/connectors/notes/validate", config=bad)
|
||||
self.assertEqual((status, body["error"]["code"]), (401, "unauthorized"))
|
||||
status, body = self.c.call("/v1/connectors/notes/list-resources", {}, cfg)
|
||||
self.assertEqual(body["output"], {"resources": [{"externalId": "inbox", "name": "Inbox", "hasChildren": True}]})
|
||||
status, body = self.c.call("/v1/connectors/notes/resolve-ancestors", {"resourceIds": ["x"]}, cfg)
|
||||
self.assertEqual(body["output"], {"ancestors": []})
|
||||
|
||||
def test_fetch_streams(self):
|
||||
status, events = self.c.stream("/v1/connectors/notes/fetch", {"mode": "incremental", "cursor": {"state": {"after": 2}}})
|
||||
self.assertEqual(status, 200)
|
||||
self.assertEqual([e["type"] for e in events], ["item", "item", "item", "checkpoint", "progress", "end"])
|
||||
self.assertEqual(events[0]["data"]["externalId"], "n2")
|
||||
self.assertEqual(base64.b64decode(events[0]["data"]["content"]), b"# Note 2")
|
||||
self.assertEqual(events[-1]["data"], {"state": {"after": 5}})
|
||||
|
||||
def test_fetch_failures(self):
|
||||
status, body = self.c.stream("/v1/connectors/notes/fetch", config={"instance": {"settings": {"fail": "before"}}})
|
||||
self.assertEqual((status, body["error"]["code"]), (401, "unauthorized"))
|
||||
status, events = self.c.stream("/v1/connectors/notes/fetch", config={"instance": {"settings": {"fail": "during"}}})
|
||||
self.assertEqual([e["type"] for e in events], ["item", "error"])
|
||||
self.assertEqual(events[1]["error"], {"code": "rate_limited", "message": "slow down", "retryable": True, "details": {"retryAfter": 7}})
|
||||
|
||||
|
||||
class BindTest(unittest.TestCase):
|
||||
def test_no_reverse_dns_on_bind(self):
|
||||
# getfqdn can stall startup for seconds where reverse DNS is slow.
|
||||
from unittest import mock
|
||||
|
||||
with mock.patch("socket.getfqdn", side_effect=AssertionError("getfqdn called")):
|
||||
server = plugin.test_server()
|
||||
server.shutdown()
|
||||
server.server_close()
|
||||
|
||||
|
||||
class WireTest(unittest.TestCase):
|
||||
def test_times(self):
|
||||
t = parse_time("2026-09-25T16:22:58.135616789Z")
|
||||
self.assertEqual(t, datetime(2026, 9, 25, 16, 22, 58, 135616, tzinfo=timezone.utc))
|
||||
self.assertEqual(to_wire(Cursor(last_sync_time=t)), {"lastSyncTime": "2026-09-25T16:22:58.135616Z"})
|
||||
naive = datetime(2026, 1, 2, 3, 4, 5)
|
||||
self.assertEqual(to_wire(naive), "2026-01-02T03:04:05Z")
|
||||
|
||||
def test_from_wire(self):
|
||||
inp = from_wire(FetchInput, {"mode": "incremental", "cursor": {"lastSyncTime": "2026-01-02T03:04:05Z", "state": {"a": 1}}, "unknown": 1})
|
||||
self.assertEqual(inp.cursor.state, {"a": 1})
|
||||
self.assertEqual(inp.cursor.last_sync_time.year, 2026)
|
||||
self.assertEqual(inp.resource_ids, [])
|
||||
|
||||
|
||||
class SignatureTest(unittest.TestCase):
|
||||
def test_sign_and_verify(self):
|
||||
now = 1_700_000_000
|
||||
sig = p.sign(b"s3cret", now, b"{}")
|
||||
p.verify_signature(b"s3cret", str(now), sig, b"{}", now=now)
|
||||
with self.assertRaisesRegex(ValueError, "bad signature"):
|
||||
p.verify_signature(b"other", str(now), sig, b"{}", now=now)
|
||||
with self.assertRaisesRegex(ValueError, "clock skew"):
|
||||
p.verify_signature(b"s3cret", str(now), sig, b"{}", now=now + 600)
|
||||
|
||||
def test_matches_go(self):
|
||||
# pluginapi.Sign([]byte("k"), 1, []byte("b")) in Go.
|
||||
import hashlib
|
||||
import hmac
|
||||
|
||||
self.assertEqual(p.sign(b"k", 1, b"b"), hmac.new(b"k", b"1.b", hashlib.sha256).hexdigest())
|
||||
|
||||
|
||||
class FakeHostAPI(BaseHTTPRequestHandler):
|
||||
store: dict = {}
|
||||
|
||||
def _reply(self, status, body=None):
|
||||
raw = b"" if body is None else json.dumps(body).encode()
|
||||
self.send_response(status)
|
||||
self.send_header("Content-Type", "application/json")
|
||||
self.send_header("Content-Length", str(len(raw)))
|
||||
self.end_headers()
|
||||
self.wfile.write(raw)
|
||||
|
||||
def _authorized(self):
|
||||
if self.headers.get("Authorization") != "Bearer tok":
|
||||
self._reply(401, {"error": {"code": "unauthorized", "message": "bad token"}})
|
||||
return False
|
||||
return True
|
||||
|
||||
def do_GET(self):
|
||||
if not self._authorized():
|
||||
return
|
||||
u = urlparse(self.path)
|
||||
q = {k: v[0] for k, v in parse_qs(u.query).items()}
|
||||
if u.path == p.HOST_KV_LIST_PATH:
|
||||
keys = sorted(k for k in self.store if k.startswith(q.get("prefix", "")) and k > q.get("after", ""))
|
||||
page = keys[:2]
|
||||
self._reply(200, {"entries": [{"key": k, "value": self.store[k], "updatedAt": "2026-01-01T00:00:00Z"} for k in page], "next": page[-1] if len(keys) > 2 else ""})
|
||||
elif q["key"] in self.store:
|
||||
self._reply(200, {"key": q["key"], "value": self.store[q["key"]], "updatedAt": "2026-01-01T00:00:00Z"})
|
||||
else:
|
||||
self._reply(404, {"error": {"code": "not_found", "message": "no key"}})
|
||||
|
||||
def do_PUT(self):
|
||||
if not self._authorized():
|
||||
return
|
||||
body = json.loads(self.rfile.read(int(self.headers["Content-Length"])))
|
||||
self.store[body["key"]] = body["value"]
|
||||
self._reply(200, {"key": body["key"], "value": body["value"]})
|
||||
|
||||
def do_DELETE(self):
|
||||
if not self._authorized():
|
||||
return
|
||||
self.store.pop(parse_qs(urlparse(self.path).query)["key"][0], None)
|
||||
self._reply(204)
|
||||
|
||||
def log_message(self, *args):
|
||||
pass
|
||||
|
||||
|
||||
class HostTest(unittest.TestCase):
|
||||
def test_kv(self):
|
||||
srv = HTTPServer(("127.0.0.1", 0), FakeHostAPI)
|
||||
threading.Thread(target=srv.serve_forever, daemon=True).start()
|
||||
url = f"http://127.0.0.1:{srv.server_address[1]}"
|
||||
try:
|
||||
h = Host(url, "tok")
|
||||
self.assertIsNone(h.kv_get("missing"))
|
||||
self.assertEqual(h.kv_get("missing", 5), 5)
|
||||
h.kv_put("a", {"n": 1}, ttl=60)
|
||||
h.kv_put("b", 2)
|
||||
h.kv_put("c", [3])
|
||||
self.assertEqual(h.kv_get("a"), {"n": 1})
|
||||
self.assertEqual([e.key for e in h.kv_items()], ["a", "b", "c"])
|
||||
h.kv_delete("a")
|
||||
self.assertIsNone(h.kv_entry("a"))
|
||||
with self.assertRaises(PluginError) as ctx:
|
||||
Host(url, "wrong").kv_get("b")
|
||||
self.assertEqual(ctx.exception.code, ErrorCode.UNAUTHORIZED)
|
||||
finally:
|
||||
srv.shutdown()
|
||||
srv.server_close()
|
||||
with self.assertRaises(PluginError) as ctx:
|
||||
Host(url, "tok", timeout=2).kv_get("b")
|
||||
self.assertEqual(ctx.exception.code, ErrorCode.UNAVAILABLE)
|
||||
|
||||
|
||||
class UnixConnection(http.client.HTTPConnection):
|
||||
def __init__(self, path):
|
||||
super().__init__("plugin", timeout=10)
|
||||
self.path = path
|
||||
|
||||
def connect(self):
|
||||
self.sock = socket.socket(socket.AF_UNIX, socket.SOCK_STREAM)
|
||||
self.sock.connect(self.path)
|
||||
|
||||
|
||||
def _free_port():
|
||||
with socket.socket() as s:
|
||||
s.bind(("127.0.0.1", 0))
|
||||
return s.getsockname()[1]
|
||||
|
||||
|
||||
@unittest.skipUnless(hasattr(socket, "AF_UNIX"), "needs unix sockets")
|
||||
class ServeTest(unittest.TestCase):
|
||||
def _start(self, env):
|
||||
proc = subprocess.Popen(
|
||||
[sys.executable, os.path.join(HERE, "fixture.py")],
|
||||
env={**{"PATH": os.environ.get("PATH", "")}, **env},
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.PIPE,
|
||||
text=True,
|
||||
)
|
||||
|
||||
def stop():
|
||||
proc.terminate()
|
||||
proc.wait(10)
|
||||
proc.stdout.close()
|
||||
proc.stderr.close()
|
||||
|
||||
self.addCleanup(stop)
|
||||
return proc
|
||||
|
||||
def test_host_mode(self):
|
||||
import tempfile
|
||||
|
||||
d = tempfile.mkdtemp()
|
||||
sock = os.path.join(d, "p.sock")
|
||||
proc = self._start({p.ENV_SOCKET: sock, p.ENV_TOKEN: "tok"})
|
||||
self.assertEqual(proc.stdout.readline().strip(), f"WEKNORA_PLUGIN|1|unix|{sock}")
|
||||
c = Client(lambda: UnixConnection(sock))
|
||||
self.assertEqual(c.request("GET", "/v1/health")[0], 401)
|
||||
status, _, _ = c.request("GET", "/v1/health", headers={"Authorization": "Bearer tok"})
|
||||
self.assertEqual(status, 200)
|
||||
proc.terminate()
|
||||
self.assertEqual(proc.wait(10), 0)
|
||||
self.assertFalse(os.path.exists(sock), "the socket is removed on shutdown")
|
||||
|
||||
def test_remote_mode(self):
|
||||
port = _free_port()
|
||||
self._start({p.ENV_ADDR: f"127.0.0.1:{port}", p.ENV_SECRET: "s3cret"})
|
||||
c = Client(lambda: http.client.HTTPConnection("127.0.0.1", port, timeout=10))
|
||||
for _ in range(100):
|
||||
try:
|
||||
c.request("GET", "/v1/health")
|
||||
break
|
||||
except OSError:
|
||||
time.sleep(0.05)
|
||||
body = json.dumps(dict(ENVELOPE, input={"query": "q"})).encode()
|
||||
ts = int(time.time())
|
||||
status, _, _ = c.request("POST", "/v1/websearch/echo/search", body, {p.TIMESTAMP_HEADER: str(ts), p.SIGNATURE_HEADER: p.sign(b"other", ts, body)})
|
||||
self.assertEqual(status, 401)
|
||||
status, _, data = c.request("POST", "/v1/websearch/echo/search", body, {p.TIMESTAMP_HEADER: str(ts), p.SIGNATURE_HEADER: p.sign(b"s3cret", ts, body)})
|
||||
self.assertEqual(status, 200, data)
|
||||
|
||||
def test_shutdown_drains_calls_in_flight(self):
|
||||
port = _free_port()
|
||||
proc = self._start({p.ENV_ADDR: f"127.0.0.1:{port}", p.ENV_SECRET: "s3cret"})
|
||||
c = Client(lambda: http.client.HTTPConnection("127.0.0.1", port, timeout=10))
|
||||
for _ in range(100):
|
||||
try:
|
||||
c.request("GET", "/v1/health")
|
||||
break
|
||||
except OSError:
|
||||
time.sleep(0.05)
|
||||
body = json.dumps(dict(ENVELOPE, input={"query": "slow"})).encode()
|
||||
ts = int(time.time())
|
||||
headers = {p.TIMESTAMP_HEADER: str(ts), p.SIGNATURE_HEADER: p.sign(b"s3cret", ts, body)}
|
||||
result = {}
|
||||
t = threading.Thread(target=lambda: result.update(r=c.request("POST", "/v1/websearch/echo/search", body, headers)))
|
||||
t.start()
|
||||
time.sleep(0.2)
|
||||
proc.terminate()
|
||||
t.join(10)
|
||||
self.assertEqual(result["r"][0], 200)
|
||||
self.assertEqual(proc.wait(10), 0)
|
||||
|
||||
def test_needs_a_secret(self):
|
||||
proc = self._start({})
|
||||
self.assertNotEqual(proc.wait(10), 0)
|
||||
self.assertIn(p.ENV_SECRET, proc.stderr.read())
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user