mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
feat(proxy): add OpenCode skills endpoint
This commit is contained in:
parent
a545c493d7
commit
5ba4a4e5f9
4 changed files with 340 additions and 0 deletions
1
litellm/proxy/opencode_endpoints/__init__.py
Normal file
1
litellm/proxy/opencode_endpoints/__init__.py
Normal file
|
|
@ -0,0 +1 @@
|
|||
"""OpenCode-compatible proxy endpoints."""
|
||||
152
litellm/proxy/opencode_endpoints/skills_endpoints.py
Normal file
152
litellm/proxy/opencode_endpoints/skills_endpoints.py
Normal file
|
|
@ -0,0 +1,152 @@
|
|||
import re
|
||||
from typing import Dict, List, Optional
|
||||
|
||||
from fastapi import Depends, FastAPI, HTTPException, Response
|
||||
|
||||
from litellm.llms.litellm_proxy.skills.handler import LiteLLMSkillsHandler
|
||||
from litellm.llms.litellm_proxy.skills.prompt_injection import (
|
||||
SkillPromptInjectionHandler,
|
||||
)
|
||||
from litellm.proxy._types import LiteLLM_SkillsTable, UserAPIKeyAuth
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
|
||||
OPENCODE_SKILLS_DEFAULT_PATH = "/opencode/skills"
|
||||
_MAX_SKILLS = 1000
|
||||
|
||||
|
||||
def _opencode_config_enabled(skills_gateway_config: Optional[dict]) -> bool:
|
||||
if not isinstance(skills_gateway_config, dict):
|
||||
return False
|
||||
opencode_config = skills_gateway_config.get("opencode", {})
|
||||
return (
|
||||
skills_gateway_config.get("enabled") is True
|
||||
and isinstance(opencode_config, dict)
|
||||
and opencode_config.get("enabled") is True
|
||||
)
|
||||
|
||||
|
||||
def _opencode_path(skills_gateway_config: Optional[dict]) -> str:
|
||||
if not isinstance(skills_gateway_config, dict):
|
||||
return OPENCODE_SKILLS_DEFAULT_PATH
|
||||
opencode_config = skills_gateway_config.get("opencode", {})
|
||||
if not isinstance(opencode_config, dict):
|
||||
return OPENCODE_SKILLS_DEFAULT_PATH
|
||||
path = opencode_config.get("path") or OPENCODE_SKILLS_DEFAULT_PATH
|
||||
path = str(path).strip() or OPENCODE_SKILLS_DEFAULT_PATH
|
||||
if not path.startswith("/"):
|
||||
path = f"/{path}"
|
||||
return path.rstrip("/") or OPENCODE_SKILLS_DEFAULT_PATH
|
||||
|
||||
|
||||
def _skill_enabled(skill: LiteLLM_SkillsTable) -> bool:
|
||||
metadata = skill.metadata if isinstance(skill.metadata, dict) else {}
|
||||
return metadata.get("enabled") is not False
|
||||
|
||||
|
||||
def _opencode_skill_name(skill: LiteLLM_SkillsTable) -> str:
|
||||
return skill.skill_id
|
||||
|
||||
|
||||
def _slug(value: str) -> str:
|
||||
slug = re.sub(r"[^a-z0-9]+", "-", value.lower()).strip("-")
|
||||
return slug or "litellm-skill"
|
||||
|
||||
|
||||
def _one_line(value: Optional[str]) -> str:
|
||||
if not value:
|
||||
return ""
|
||||
return " ".join(str(value).split())
|
||||
|
||||
|
||||
def _generated_skill_md(skill: LiteLLM_SkillsTable) -> bytes:
|
||||
title = skill.display_title or skill.skill_id
|
||||
description = _one_line(skill.description or skill.instructions or title)
|
||||
instructions = (skill.instructions or skill.description or title).strip()
|
||||
return (
|
||||
"---\n"
|
||||
f"name: {_slug(skill.skill_id)}\n"
|
||||
f"description: {description}\n"
|
||||
"---\n\n"
|
||||
f"# {title}\n\n"
|
||||
f"{instructions}\n"
|
||||
).encode("utf-8")
|
||||
|
||||
|
||||
def _skill_files(skill: LiteLLM_SkillsTable) -> Dict[str, bytes]:
|
||||
files = SkillPromptInjectionHandler().extract_all_files(skill)
|
||||
if "SKILL.md" not in files:
|
||||
files["SKILL.md"] = _generated_skill_md(skill)
|
||||
return files
|
||||
|
||||
|
||||
async def _enabled_skills(
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
) -> List[LiteLLM_SkillsTable]:
|
||||
skills = await LiteLLMSkillsHandler.list_skills(
|
||||
limit=_MAX_SKILLS,
|
||||
offset=0,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
return [skill for skill in skills if _skill_enabled(skill)]
|
||||
|
||||
|
||||
def _sorted_files(files: Dict[str, bytes]) -> List[str]:
|
||||
return sorted(files)
|
||||
|
||||
|
||||
async def opencode_skills_index(
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
skills = await _enabled_skills(user_api_key_dict)
|
||||
return {
|
||||
"skills": [
|
||||
{
|
||||
"name": _opencode_skill_name(skill),
|
||||
"files": _sorted_files(_skill_files(skill)),
|
||||
}
|
||||
for skill in skills
|
||||
]
|
||||
}
|
||||
|
||||
|
||||
async def opencode_skill_file(
|
||||
skill_name: str,
|
||||
file_path: str,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
for skill in await _enabled_skills(user_api_key_dict):
|
||||
if _opencode_skill_name(skill) != skill_name:
|
||||
continue
|
||||
files = _skill_files(skill)
|
||||
content = files.get(file_path)
|
||||
if content is None:
|
||||
break
|
||||
media_type = (
|
||||
"text/markdown; charset=utf-8"
|
||||
if file_path.endswith(".md")
|
||||
else "application/octet-stream"
|
||||
)
|
||||
return Response(content=content, media_type=media_type)
|
||||
raise HTTPException(status_code=404, detail="Skill file not found")
|
||||
|
||||
|
||||
def _add_route(app: FastAPI, path: str, endpoint):
|
||||
for route in app.routes:
|
||||
if getattr(route, "path", None) == path and "GET" in getattr(
|
||||
route, "methods", set()
|
||||
):
|
||||
return
|
||||
app.add_api_route(path=path, endpoint=endpoint, methods=["GET"])
|
||||
|
||||
|
||||
def initialize_opencode_skills_endpoint(
|
||||
app: FastAPI,
|
||||
skills_gateway_config: Optional[dict],
|
||||
) -> None:
|
||||
if not _opencode_config_enabled(skills_gateway_config):
|
||||
return
|
||||
|
||||
path = _opencode_path(skills_gateway_config)
|
||||
_add_route(app, path, opencode_skills_index)
|
||||
_add_route(app, f"{path}/index.json", opencode_skills_index)
|
||||
_add_route(app, f"{path}/{{skill_name}}/{{file_path:path}}", opencode_skill_file)
|
||||
|
|
@ -4551,6 +4551,15 @@ class ProxyConfig:
|
|||
general_settings = config.get("general_settings", {})
|
||||
if general_settings is None:
|
||||
general_settings = {}
|
||||
from litellm.proxy.opencode_endpoints.skills_endpoints import (
|
||||
initialize_opencode_skills_endpoint,
|
||||
)
|
||||
|
||||
initialize_opencode_skills_endpoint(
|
||||
app=app,
|
||||
skills_gateway_config=config.get("skills_gateway"),
|
||||
)
|
||||
|
||||
_enable_hc_routing = False
|
||||
_hc_staleness = None
|
||||
_hc_ignore_transient = False
|
||||
|
|
|
|||
178
tests/test_litellm/proxy/test_opencode_skills_endpoints.py
Normal file
178
tests/test_litellm/proxy/test_opencode_skills_endpoints.py
Normal file
|
|
@ -0,0 +1,178 @@
|
|||
from unittest.mock import AsyncMock
|
||||
from io import BytesIO
|
||||
from zipfile import ZipFile
|
||||
|
||||
from fastapi import FastAPI
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from litellm.proxy._types import LiteLLM_SkillsTable, UserAPIKeyAuth
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
|
||||
|
||||
def _client(config: dict, auth: UserAPIKeyAuth | None = None) -> TestClient:
|
||||
from litellm.proxy.opencode_endpoints.skills_endpoints import (
|
||||
initialize_opencode_skills_endpoint,
|
||||
)
|
||||
|
||||
app = FastAPI()
|
||||
if auth is not None:
|
||||
app.dependency_overrides[user_api_key_auth] = lambda: auth
|
||||
initialize_opencode_skills_endpoint(app=app, skills_gateway_config=config)
|
||||
return TestClient(app)
|
||||
|
||||
|
||||
def test_should_not_register_opencode_skills_endpoint_when_disabled():
|
||||
client = _client({})
|
||||
|
||||
response = client.get("/opencode/skills")
|
||||
|
||||
assert response.status_code == 404
|
||||
|
||||
|
||||
def test_should_return_opencode_skills_index_when_enabled(monkeypatch):
|
||||
from litellm.llms.litellm_proxy.skills.handler import LiteLLMSkillsHandler
|
||||
|
||||
list_skills = AsyncMock(
|
||||
return_value=[
|
||||
LiteLLM_SkillsTable(
|
||||
skill_id="litellm_skill_writer",
|
||||
display_title="Writer",
|
||||
description="Draft clean release notes",
|
||||
instructions="Use this skill to draft release notes.",
|
||||
)
|
||||
]
|
||||
)
|
||||
monkeypatch.setattr(LiteLLMSkillsHandler, "list_skills", list_skills)
|
||||
auth = UserAPIKeyAuth(user_id="user-1")
|
||||
client = _client(
|
||||
{"enabled": True, "opencode": {"enabled": True}},
|
||||
auth=auth,
|
||||
)
|
||||
|
||||
response = client.get("/opencode/skills/index.json")
|
||||
|
||||
assert response.status_code == 200
|
||||
assert response.json() == {
|
||||
"skills": [{"name": "litellm_skill_writer", "files": ["SKILL.md"]}]
|
||||
}
|
||||
list_skills.assert_awaited_once_with(limit=1000, offset=0, user_api_key_dict=auth)
|
||||
|
||||
|
||||
def test_should_omit_disabled_litellm_skills(monkeypatch):
|
||||
from litellm.llms.litellm_proxy.skills.handler import LiteLLMSkillsHandler
|
||||
|
||||
monkeypatch.setattr(
|
||||
LiteLLMSkillsHandler,
|
||||
"list_skills",
|
||||
AsyncMock(
|
||||
return_value=[
|
||||
LiteLLM_SkillsTable(
|
||||
skill_id="litellm_skill_enabled",
|
||||
metadata={"enabled": True},
|
||||
),
|
||||
LiteLLM_SkillsTable(
|
||||
skill_id="litellm_skill_disabled",
|
||||
metadata={"enabled": False},
|
||||
),
|
||||
]
|
||||
),
|
||||
)
|
||||
client = _client(
|
||||
{"enabled": True, "opencode": {"enabled": True}},
|
||||
auth=UserAPIKeyAuth(user_id="user-1"),
|
||||
)
|
||||
|
||||
response = client.get("/opencode/skills")
|
||||
|
||||
assert response.status_code == 200
|
||||
assert response.json() == {
|
||||
"skills": [{"name": "litellm_skill_enabled", "files": ["SKILL.md"]}]
|
||||
}
|
||||
|
||||
|
||||
def test_should_serve_generated_skill_markdown(monkeypatch):
|
||||
from litellm.llms.litellm_proxy.skills.handler import LiteLLMSkillsHandler
|
||||
|
||||
monkeypatch.setattr(
|
||||
LiteLLMSkillsHandler,
|
||||
"list_skills",
|
||||
AsyncMock(
|
||||
return_value=[
|
||||
LiteLLM_SkillsTable(
|
||||
skill_id="litellm_skill_writer",
|
||||
display_title="Writer",
|
||||
description="Draft clean release notes",
|
||||
instructions="Use this skill to draft release notes.",
|
||||
)
|
||||
]
|
||||
),
|
||||
)
|
||||
client = _client(
|
||||
{"enabled": True, "opencode": {"enabled": True}},
|
||||
auth=UserAPIKeyAuth(user_id="user-1"),
|
||||
)
|
||||
|
||||
response = client.get("/opencode/skills/litellm_skill_writer/SKILL.md")
|
||||
|
||||
assert response.status_code == 200
|
||||
assert response.text == (
|
||||
"---\n"
|
||||
"name: litellm-skill-writer\n"
|
||||
"description: Draft clean release notes\n"
|
||||
"---\n\n"
|
||||
"# Writer\n\n"
|
||||
"Use this skill to draft release notes.\n"
|
||||
)
|
||||
|
||||
|
||||
def test_should_serve_zip_backed_skill_files_from_custom_path(monkeypatch):
|
||||
from litellm.llms.litellm_proxy.skills.handler import LiteLLMSkillsHandler
|
||||
|
||||
zip_buffer = BytesIO()
|
||||
with ZipFile(zip_buffer, "w") as zip_file:
|
||||
zip_file.writestr(
|
||||
"writer/SKILL.md",
|
||||
"---\nname: writer\ndescription: Draft notes\n---\n\nUse me.",
|
||||
)
|
||||
zip_file.writestr("writer/references/example.txt", "example")
|
||||
|
||||
monkeypatch.setattr(
|
||||
LiteLLMSkillsHandler,
|
||||
"list_skills",
|
||||
AsyncMock(
|
||||
return_value=[
|
||||
LiteLLM_SkillsTable(
|
||||
skill_id="litellm_skill_writer",
|
||||
file_content=zip_buffer.getvalue(),
|
||||
file_name="writer.zip",
|
||||
)
|
||||
]
|
||||
),
|
||||
)
|
||||
client = _client(
|
||||
{
|
||||
"enabled": True,
|
||||
"opencode": {"enabled": True, "path": "/custom/skills"},
|
||||
},
|
||||
auth=UserAPIKeyAuth(user_id="user-1"),
|
||||
)
|
||||
|
||||
index = client.get("/custom/skills/index.json")
|
||||
skill_md = client.get("/custom/skills/litellm_skill_writer/SKILL.md")
|
||||
reference = client.get("/custom/skills/litellm_skill_writer/references/example.txt")
|
||||
|
||||
assert index.status_code == 200
|
||||
assert index.json() == {
|
||||
"skills": [
|
||||
{
|
||||
"name": "litellm_skill_writer",
|
||||
"files": ["SKILL.md", "references/example.txt"],
|
||||
}
|
||||
]
|
||||
}
|
||||
assert skill_md.status_code == 200
|
||||
assert (
|
||||
skill_md.text == "---\nname: writer\ndescription: Draft notes\n---\n\nUse me."
|
||||
)
|
||||
assert reference.status_code == 200
|
||||
assert reference.text == "example"
|
||||
Loading…
Add table
Reference in a new issue