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:
lyingbug
2026-09-26 14:46:03 +08:00
committed by GitHub
parent ff49f9491d
commit 728007ada9
52 changed files with 3761 additions and 101 deletions
+34 -1
View File
@@ -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
+106
View File
@@ -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"![Output of cell {cell + 1}]({ref})")
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()
+18
View File
@@ -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"
+22
View File
@@ -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: 每段文本输出保留的字符数 }
+73
View File
@@ -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("![Output of cell 2](images/cell1-1.png)", 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()
+28 -7
View File
@@ -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)}`)
}
+22 -2
View File
@@ -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: {
+22 -2
View File
@@ -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: {
+22 -2
View File
@@ -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: {
+22 -2
View File
@@ -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: {
+22 -2
View File
@@ -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: {
+22 -1
View File
@@ -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)
+9 -5
View File
@@ -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)
}
}
+7 -4
View File
@@ -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,
+35 -7
View File
@@ -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)
+79 -13
View File
@@ -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 删除插件及其全部版本;各空间的开关和配置保留,重新安装后恢复
+16 -3
View File
@@ -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)
+11
View File
@@ -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)
}
}
+3 -2
View File
@@ -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) {
+92 -18
View File
@@ -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)
}
}
}
+130
View File
@@ -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")
}
}
+132 -13
View File
@@ -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.
+100 -2
View File
@@ -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)
}
}
+64 -8
View File
@@ -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)
}
}
+248
View File
@@ -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()
}
}
+196
View File
@@ -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)
}
}
+2
View File
@@ -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)
}
+6 -2
View File
@@ -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
View File
@@ -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.
+175
View File
@@ -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)
},
}))
}
+3
View File
@@ -0,0 +1,3 @@
build/
dist/
*.egg-info/
+133
View File
@@ -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 `![](ref)` 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
```
+23
View File
@@ -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)."""
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 = ""
+78
View File
@@ -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![](img/dot.png)",
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()
+363
View File
@@ -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()