Add qualifier prompt test page

This commit is contained in:
Your Name
2026-06-15 22:56:40 +05:00
parent 2e4cbee4a5
commit d306ca42a1
3 changed files with 298 additions and 1 deletions
+182
View File
@@ -1,5 +1,6 @@
from __future__ import annotations
import asyncio
import json
import base64
import hashlib
@@ -26,6 +27,7 @@ from .db import fetch_setting, get_pool
from .security import hash_password, new_token, token_hash, verify_password
from .text_utils import normalize_hash_tag, parse_categories
from .vk_api import VKAPIClient, normalize_vk_source
from .workers.ai_qualifier import normalize_model, response_usage
COOKIE_NAME = "vk_parser_admin"
VK_OAUTH_VERIFIER_COOKIE = "vk_oauth_verifier"
@@ -914,6 +916,71 @@ async def category_rows() -> list[dict[str, Any]]:
return result
async def prompt_test_posts(selected_post_id: int | None = None) -> tuple[list[dict[str, Any]], dict[str, Any] | None]:
pool = await get_pool()
rows = await pool.fetch(
"""
SELECT rp.id, rp.raw_text, rp.original_url, rp.posted_at, rp.created_at, rp.qualification_status,
rp.qualification_score, s.name AS source_name, s.tag AS source_tag,
(SELECT COUNT(*) FROM raw_post_media m WHERE m.raw_post_id=rp.id) AS media_count,
(SELECT ARRAY_AGG(DISTINCT m.media_type ORDER BY m.media_type)
FROM raw_post_media m WHERE m.raw_post_id=rp.id) AS media_types
FROM raw_posts rp
JOIN sources s ON s.id=rp.source_id
ORDER BY rp.created_at DESC, rp.id DESC
LIMIT 150
"""
)
posts = []
for row in rows:
item = dict(row)
text = str(item.get("raw_text") or "").strip()
item["created_at_fmt"] = format_dt(item.get("created_at"))
item["posted_at_fmt"] = format_dt(item.get("posted_at"))
item["label"] = f"#{item['id']} · {item.get('source_name') or 'source'} · {item['created_at_fmt']}"
item["snippet"] = text[:180] + ("..." if len(text) > 180 else "")
posts.append(item)
selected = next((item for item in posts if int(item["id"]) == int(selected_post_id or 0)), None)
if not selected and selected_post_id:
row = await pool.fetchrow(
"""
SELECT rp.id, rp.raw_text, rp.original_url, rp.posted_at, rp.created_at, rp.qualification_status,
rp.qualification_score, s.name AS source_name, s.tag AS source_tag,
(SELECT COUNT(*) FROM raw_post_media m WHERE m.raw_post_id=rp.id) AS media_count,
(SELECT ARRAY_AGG(DISTINCT m.media_type ORDER BY m.media_type)
FROM raw_post_media m WHERE m.raw_post_id=rp.id) AS media_types
FROM raw_posts rp
JOIN sources s ON s.id=rp.source_id
WHERE rp.id=$1
""",
selected_post_id,
)
if row:
selected = dict(row)
text = str(selected.get("raw_text") or "").strip()
selected["created_at_fmt"] = format_dt(selected.get("created_at"))
selected["posted_at_fmt"] = format_dt(selected.get("posted_at"))
selected["label"] = f"#{selected['id']} · {selected.get('source_name') or 'source'} · {selected['created_at_fmt']}"
selected["snippet"] = text[:180] + ("..." if len(text) > 180 else "")
posts.insert(0, selected)
if not selected and posts:
selected = posts[0]
return posts, selected
def prompt_test_payload(post: dict[str, Any], max_text_chars: int) -> list[dict[str, Any]]:
return [
{
"id": int(post["id"]),
"source": post.get("source_name") or "",
"original_url": post.get("original_url") or "",
"media_count": int(post.get("media_count") or 0),
"media_types": list(post.get("media_types") or []),
"text": str(post.get("raw_text") or "")[:max_text_chars],
}
]
def setting_value(settings_rows: list[dict], key: str, default: Any = None) -> Any:
for row in settings_rows:
if row.get("key") == key:
@@ -2195,6 +2262,121 @@ async def raw_post_detail(request: Request, post_id: int):
return templates.TemplateResponse("raw_post_detail.html", base_context(request, user, post=post))
@app.get("/prompt-test", response_class=HTMLResponse)
async def prompt_test(request: Request, post_id: int | None = None):
user = await get_current_user(request)
if not user:
return redirect("/login")
posts, selected = await prompt_test_posts(post_id)
max_text_chars = max(100, int(await fetch_setting("ai_qualifier_max_text_chars", 2000) or 2000))
prompt = str(await fetch_setting("ai_qualifier_prompt", "") or "")
provider = str(await fetch_setting("ai_qualifier_provider", "openrouter") or "openrouter")
model = str(await fetch_setting("ai_qualifier_model", "") or "")
payload = prompt_test_payload(selected, max_text_chars) if selected else []
return templates.TemplateResponse(
"prompt_test.html",
base_context(
request,
user,
posts=posts,
selected=selected,
prompt=prompt,
provider=provider,
model=model,
max_text_chars=max_text_chars,
payload_json=json.dumps(payload, ensure_ascii=False, indent=2),
result=None,
error=None,
),
)
@app.post("/prompt-test", response_class=HTMLResponse)
async def prompt_test_run(
request: Request,
csrf_token: str = Form(...),
post_id: int = Form(...),
prompt: str = Form(...),
):
user = await get_current_user(request)
if not user:
return redirect("/login")
require_csrf(user, csrf_token)
posts, selected = await prompt_test_posts(post_id)
max_text_chars = max(100, int(await fetch_setting("ai_qualifier_max_text_chars", 2000) or 2000))
provider = str(await fetch_setting("ai_qualifier_provider", "openrouter") or "openrouter")
model = str(await fetch_setting("ai_qualifier_model", "") or "").strip()
api_key = str(await fetch_setting("ai_qualifier_api_key", "") or "").strip()
api_base = str(await fetch_setting("ai_qualifier_api_base", "") or "").strip()
temperature = max(0.0, float(await fetch_setting("ai_qualifier_temperature", 0.0) or 0.0))
timeout = max(10, int(await fetch_setting("ai_qualifier_timeout_sec", 120) or 120))
normalized_prompt = prompt.replace("\\r\\n", "\n").replace("\\n", "\n").strip()
payload = prompt_test_payload(selected, max_text_chars) if selected else []
result = None
error = None
if not selected:
error = "Пост не найден."
elif not model or not api_key:
error = "Не заполнены модель или API-ключ квалификатора в настройках."
elif not normalized_prompt:
error = "Промпт пустой."
else:
try:
import litellm
kwargs: dict[str, Any] = {
"model": normalize_model(provider, model),
"messages": [
{"role": "system", "content": normalized_prompt},
{"role": "user", "content": json.dumps(payload, ensure_ascii=False)},
],
"temperature": temperature,
"timeout": timeout,
}
if api_key:
kwargs["api_key"] = api_key
if api_base:
kwargs["api_base"] = api_base
response = await asyncio.to_thread(litellm.completion, **kwargs)
content = response.choices[0].message.content
parsed = None
raw_json = str(content or "").strip()
try:
if raw_json.startswith("```"):
raw_json = raw_json.strip("`")
if raw_json.lower().startswith("json"):
raw_json = raw_json[4:].strip()
parsed = json.loads(raw_json)
except Exception:
parsed = None
result = {
"content": content,
"parsed_json": json.dumps(parsed, ensure_ascii=False, indent=2) if parsed is not None else "",
"usage": response_usage(response),
"model": kwargs["model"],
}
await audit(user["id"], "prompt_test.run", "raw_post", int(selected["id"]), {"provider": provider, "model": model})
except Exception as exc:
logger.exception("Prompt test failed")
error = str(exc)
return templates.TemplateResponse(
"prompt_test.html",
base_context(
request,
user,
posts=posts,
selected=selected,
prompt=normalized_prompt,
provider=provider,
model=model,
max_text_chars=max_text_chars,
payload_json=json.dumps(payload, ensure_ascii=False, indent=2),
result=result,
error=error,
),
)
@app.get("/workers", response_class=HTMLResponse)
async def workers(request: Request):
user = await get_current_user(request)