Add qualifier prompt test page
This commit is contained in:
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user