From a7c1c37f014c40a692eadf6f4e51d2f116686a60 Mon Sep 17 00:00:00 2001 From: mamad Date: Fri, 28 Aug 2026 19:34:00 +0330 Subject: [PATCH] core: add provider integrations, reasoning controls, and multimodal vision support --- agy_bridge.py | 133 ++++++++++++ core/llm.py | 445 ++++++++++++++++++++++++++++++++++++--- services/ai_processor.py | 286 +++++++++++++++++++++++-- tests/test_prompts.py | 209 ++++++++++++++++++ 4 files changed, 1019 insertions(+), 54 deletions(-) create mode 100644 agy_bridge.py create mode 100644 tests/test_prompts.py diff --git a/agy_bridge.py b/agy_bridge.py new file mode 100644 index 0000000..0a24c87 --- /dev/null +++ b/agy_bridge.py @@ -0,0 +1,133 @@ +#!/usr/bin/env python3 +""" +Lightweight local HTTP bridge for AGY (Antigravity CLI). +Exposes an OpenAI-compatible /v1/chat/completions endpoint on port 8088 +so Docker containers can seamlessly send AI requests to the host's agy CLI. +""" + +import sys +import os +import json +import shutil +import logging +import subprocess +from http.server import ThreadingHTTPServer, BaseHTTPRequestHandler + +logging.basicConfig(level=logging.INFO, format="%(asctime)s [%(levelname)s] agy_bridge: %(message)s") +logger = logging.getLogger("agy_bridge") + +AGY_BIN = shutil.which("agy") or ("/home/mamad/.local/bin/agy" if os.path.exists("/home/mamad/.local/bin/agy") else "agy") +PORT = int(os.getenv("AGY_BRIDGE_PORT", "8088")) +HOST = "0.0.0.0" + + +class AGYBridgeHandler(BaseHTTPRequestHandler): + def do_GET(self): + if self.path in ("/", "/health", "/v1/models"): + resp = json.dumps({"status": "ok", "service": "agy_bridge", "agy_bin": AGY_BIN}).encode("utf-8") + self.send_response(200) + self.send_header("Content-Type", "application/json") + self.send_header("Content-Length", str(len(resp))) + self.end_headers() + self.wfile.write(resp) + else: + self.send_error(404, "Not Found") + + def do_POST(self): + if not self.path.startswith("/v1/chat/completions") and self.path != "/chat/completions": + self.send_error(404, "Not Found") + return + + try: + content_length = int(self.headers.get("Content-Length", 0)) + body = self.rfile.read(content_length).decode("utf-8") + data = json.loads(body) + + messages = data.get("messages", []) + effort = data.get("reasoning_effort", "") + + system_instructions = [] + user_prompts = [] + for msg in messages: + role = msg.get("role", "user") + content = msg.get("content", "") + if role == "system": + system_instructions.append(content) + else: + user_prompts.append(content) + + full_prompt = "" + if system_instructions: + full_prompt += f"System Instruction:\n" + "\n".join(system_instructions) + "\n\n" + full_prompt += "User Prompt:\n" + "\n".join(user_prompts) + + cmd = [AGY_BIN, "-p", full_prompt, "--output-format", "json"] + if effort in ("low", "medium", "high"): + cmd.extend(["--effort", effort]) + + logger.info(f"Executing agy for prompt length: {len(full_prompt)} chars") + proc = subprocess.run( + cmd, + capture_output=True, + text=True, + timeout=120, + env=os.environ.copy() + ) + + if proc.returncode != 0: + logger.error(f"agy error (code {proc.returncode}): {proc.stderr}") + self.send_error(500, f"agy execution error: {proc.stderr}") + return + + try: + res_data = json.loads(proc.stdout) + content = res_data.get("response", "") + except Exception: + content = proc.stdout + + resp_payload = { + "id": "agy-response", + "object": "chat.completion", + "created": 1234567890, + "model": data.get("model", "antigravity"), + "choices": [ + { + "index": 0, + "message": { + "role": "assistant", + "content": content + }, + "finish_reason": "stop" + } + ] + } + resp_bytes = json.dumps(resp_payload, ensure_ascii=False).encode("utf-8") + + self.send_response(200) + self.send_header("Content-Type", "application/json") + self.send_header("Content-Length", str(len(resp_bytes))) + self.end_headers() + self.wfile.write(resp_bytes) + logger.info("Successfully handled AGY chat completion request.") + + except Exception as e: + logger.error(f"Failed to handle request: {e}", exc_info=True) + self.send_error(500, str(e)) + + def log_message(self, format, *args): + # Override default noisy stderr logging + pass + + +def main(): + logger.info(f"Starting AGY Bridge Server on http://{HOST}:{PORT} (AGY: {AGY_BIN})") + server = ThreadingHTTPServer((HOST, PORT), AGYBridgeHandler) + try: + server.serve_forever() + except KeyboardInterrupt: + logger.info("Shutting down AGY Bridge Server...") + server.shutdown() + + +if __name__ == "__main__": + main() diff --git a/core/llm.py b/core/llm.py index 1b2a2f0..68148af 100644 --- a/core/llm.py +++ b/core/llm.py @@ -1,10 +1,20 @@ import os import json import time +import asyncio +import shutil +import base64 import httpx import logging -from typing import Dict, Any, Optional +from typing import Dict, Any, Optional, List, Callable, Awaitable from core.metrics import AI_REQUESTS_TOTAL, AI_LATENCY_SECONDS +from db.models import AIProviderProfile + +try: + from google.antigravity import Agent as AgyAgent, LocalAgentConfig as AgyLocalAgentConfig + HAS_AGY_SDK = True +except ImportError: + HAS_AGY_SDK = False logger = logging.getLogger(__name__) @@ -15,43 +25,307 @@ class LLMClient: api_key: Optional[str] = None, model: Optional[str] = None, base_url: Optional[str] = None, + reasoning_effort: Optional[str] = None, + repo: Optional[Any] = None, + on_fallback_alert: Optional[Callable[[AIProviderProfile, AIProviderProfile, str, int], Awaitable[None]]] = None, + on_chain_failure_alert: Optional[Callable[[List[AIProviderProfile], str], Awaitable[None]]] = None, ): + self.repo = repo self.provider = provider or os.getenv("AI_PROVIDER", "openai").lower() self.api_key = api_key or os.getenv("AI_API_KEY", "") - self.model = model or os.getenv("AI_MODEL", "orcarouter/auto" if self.provider == "openai" else "gemini-1.5-flash") - self.base_url = base_url or os.getenv("AI_BASE_URL", "https://api.orcarouter.ai/v1") + default_model = "antigravity" if self.provider == "agy" else ("orcarouter/auto" if self.provider == "openai" else "gemini-1.5-flash") + self.model = model or os.getenv("AI_MODEL", default_model) + default_base_url = "http://localhost:8000/v1" if self.provider == "agy" else "https://api.orcarouter.ai/v1" + self.base_url = base_url or os.getenv("AI_BASE_URL", default_base_url) + self.reasoning_effort = (reasoning_effort or os.getenv("AI_REASONING_EFFORT", "")).strip().lower() + fallbacks = [m.strip() for m in os.getenv("AI_FALLBACK_MODELS", "").split(",") if m.strip()] + self.models = [self.model] + [m for m in fallbacks if m != self.model] + self.max_retries_per_model = int(os.getenv("AI_MAX_RETRIES", "2")) + self.retry_backoff_seconds = float(os.getenv("AI_RETRY_BACKOFF_SECONDS", "2")) + self.last_used_model = self.model + self.last_used_provider = self.provider + self.on_fallback_alert = on_fallback_alert + self.on_chain_failure_alert = on_chain_failure_alert - async def generate_json(self, prompt: str, system_prompt: Optional[str] = None, action_name: str = "general") -> Dict[str, Any]: - """Send prompt to LLM and parse JSON response.""" + async def sync_config_from_repo(self): + if not self.repo: + return + try: + active_profile = await self.repo.get_active_provider_profile() + if active_profile: + self.provider = active_profile.provider_type.strip().lower() + self.model = active_profile.model.strip() + self.base_url = active_profile.base_url.strip() + self.api_key = active_profile.api_key.strip() + self.reasoning_effort = active_profile.reasoning_effort.strip().lower() + else: + db_provider = await self.repo.get_setting("ai_provider") + if db_provider: + self.provider = db_provider.strip().lower() + + db_model = await self.repo.get_setting("ai_model") + if db_model: + self.model = db_model.strip() + + db_base_url = await self.repo.get_setting("ai_base_url") + if db_base_url is not None: + self.base_url = db_base_url.strip() + + db_api_key = await self.repo.get_setting("ai_api_key") + if db_api_key is not None: + self.api_key = db_api_key.strip() + + db_reasoning = await self.repo.get_setting("ai_reasoning_effort") + if db_reasoning is not None: + self.reasoning_effort = db_reasoning.strip().lower() + + fallbacks = [m.strip() for m in os.getenv("AI_FALLBACK_MODELS", "").split(",") if m.strip()] + self.models = [self.model] + [m for m in fallbacks if m != self.model] + except Exception as e: + logger.debug(f"Could not sync AI config from DB: {e}") + + async def get_fallback_chain(self) -> List[AIProviderProfile]: + """Return the ordered list of AI provider profiles manually chained via fallback_provider_id.""" + if self.repo: + try: + profiles = await self.repo.get_provider_profiles() + if profiles: + profile_map = {p.id: p for p in profiles if p.id is not None} + active = next((p for p in profiles if p.is_active), profiles[0]) + chain = [active] + visited = {active.id} + curr = active + while curr.fallback_provider_id and curr.fallback_provider_id in profile_map: + next_p = profile_map[curr.fallback_provider_id] + if next_p.id in visited: + break # prevent infinite cycle + chain.append(next_p) + visited.add(next_p.id) + curr = next_p + return chain + except Exception as e: + logger.debug(f"Error fetching provider profiles from repo: {e}") + + # Fallback to current memory/env configuration + return [ + AIProviderProfile( + id=0, + name="Default Provider", + provider_type=self.provider, + model=self.model, + base_url=self.base_url, + api_key=self.api_key, + reasoning_effort=self.reasoning_effort, + is_active=True + ) + ] + + async def generate_json( + self, + prompt: str, + system_prompt: Optional[str] = None, + action_name: str = "general", + image_path: Optional[str] = None + ) -> Dict[str, Any]: + """Send prompt (and optional image) to LLM, falling back across provider chain and notifying admins on failure.""" + await self.sync_config_from_repo() + chain = await self.get_fallback_chain() start_time = time.time() status = "error" + last_error: Optional[Exception] = None + result_text: Optional[str] = None + successful_profile: Optional[AIProviderProfile] = None + try: - if self.provider == "gemini" and "orcarouter" not in (self.base_url or ""): - result = await self._call_gemini(prompt, system_prompt) - else: - result = await self._call_openai(prompt, system_prompt) - status = "success" - return result - except Exception as e: - logger.error(f"LLM generation failed ({self.provider}/{self.model}): {e}") - raise + for i, profile in enumerate(chain): + p_type = profile.provider_type.strip().lower() + p_model = profile.model.strip() + p_base_url = profile.base_url.strip() + p_key = profile.api_key.strip() + p_effort = profile.reasoning_effort.strip().lower() + + for attempt in range(self.max_retries_per_model + 1): + try: + if p_type == "gemini" and "orcarouter" not in p_base_url: + result = await self._call_gemini( + prompt=prompt, + system_prompt=system_prompt, + model=p_model, + api_key=p_key or self.api_key, + reasoning_effort=p_effort or self.reasoning_effort, + image_path=image_path + ) + elif p_type == "agy": + if not os.path.exists("/.dockerenv") and (shutil.which("agy") or os.path.exists("/home/mamad/.local/bin/agy")): + result = await self._call_agy_cli( + prompt=prompt, + system_prompt=system_prompt, + effort=p_effort or self.reasoning_effort, + image_path=image_path + ) + else: + result = await self._call_openai( + prompt=prompt, + system_prompt=system_prompt, + model=p_model, + base_url=p_base_url, + api_key=p_key, + reasoning_effort=p_effort or self.reasoning_effort, + is_agy=True, + image_path=image_path + ) + else: + result = await self._call_openai( + prompt=prompt, + system_prompt=system_prompt, + model=p_model, + base_url=p_base_url, + api_key=p_key or self.api_key, + reasoning_effort=p_effort or self.reasoning_effort, + is_agy=False, + image_path=image_path + ) + + status = "success" + self.last_used_model = p_model + self.last_used_provider = p_type + successful_profile = profile + result_text = json.dumps(result, ensure_ascii=False) if isinstance(result, dict) else str(result) + return result + + except Exception as e: + last_error = e + retryable = self._is_retryable(e) + logger.warning( + f"LLM call failed (provider={p_type}, model={p_model}, attempt={attempt + 1}/{self.max_retries_per_model + 1}, retryable={retryable}): {e}" + ) + if not retryable or attempt == self.max_retries_per_model: + break + await asyncio.sleep(self.retry_backoff_seconds * (attempt + 1)) + + # Profile failed across all retries; trigger fallback alert if more providers exist + if i + 1 < len(chain): + next_profile = chain[i + 1] + logger.warning( + f"AI Provider '{profile.name}' ({p_type}/{p_model}) failed: {last_error}. Falling back to '{next_profile.name}' ({next_profile.provider_type}/{next_profile.model}) (step {i + 1})." + ) + if self.on_fallback_alert: + try: + await self.on_fallback_alert( + profile, + next_profile, + str(last_error), + i + 1 + ) + except Exception as alert_err: + logger.error(f"Error executing AI fallback alert callback: {alert_err}") + + # All providers exhausted + if self.on_chain_failure_alert and len(chain) > 1: + try: + await self.on_chain_failure_alert( + chain, + str(last_error) + ) + except Exception as alert_err: + logger.error(f"Error executing AI chain failure alert callback: {alert_err}") + + logger.error(f"LLM generation failed across all providers in chain: {last_error}") + raise last_error if last_error else RuntimeError("LLM generation failed across all providers in chain") + finally: duration = time.time() - start_time AI_LATENCY_SECONDS.labels(action=action_name).observe(duration) AI_REQUESTS_TOTAL.labels(action=action_name, status=status).inc() + if self.repo: + try: + used_prov = successful_profile.provider_type if successful_profile else self.last_used_provider + used_mod = successful_profile.model if successful_profile else (self.last_used_model or self.model) + await self.repo.record_ai_log( + action_name=action_name, + provider=used_prov, + model=used_mod, + prompt=prompt, + system_prompt=system_prompt, + response_text=result_text, + duration_sec=duration, + status=status, + error_message=str(last_error) if last_error and status != "success" else None, + ) + except Exception as log_err: + logger.error(f"Failed to record AI log in database: {log_err}", exc_info=True) + + @staticmethod + def _is_retryable(error: Exception) -> bool: + """Rate limits, upstream outages and transport errors are worth another shot.""" + if isinstance(error, httpx.HTTPStatusError): + return error.response.status_code in (408, 409, 425, 429, 500, 502, 503, 504) + return isinstance(error, (httpx.TransportError, json.JSONDecodeError, KeyError, ValueError)) + + @staticmethod + def _parse_json_content(content: Optional[str]) -> Dict[str, Any]: + """Parse a model reply that may be empty or wrapped in a markdown code fence.""" + text = (content or "").strip() + if not text: + raise ValueError("model returned an empty response") + + if text.startswith("```"): + text = text.split("\n", 1)[-1] if "\n" in text else text + text = text.rsplit("```", 1)[0].strip() + + try: + return json.loads(text) + except json.JSONDecodeError: + # Some models prepend prose before the JSON object. + start, end = text.find("{"), text.rfind("}") + if start != -1 and end > start: + return json.loads(text[start:end + 1]) + raise + + async def _call_gemini( + self, + prompt: str, + system_prompt: Optional[str] = None, + model: Optional[str] = None, + api_key: Optional[str] = None, + reasoning_effort: Optional[str] = None, + image_path: Optional[str] = None, + ) -> Dict[str, Any]: + model = model or self.model + key = (api_key or self.api_key or "").strip() + eff = (reasoning_effort or self.reasoning_effort or "").strip().lower() + url = f"https://generativelanguage.googleapis.com/v1beta/models/{model}:generateContent?key={key}" + gen_config: Dict[str, Any] = { + "responseMimeType": "application/json", + "temperature": 0.2, + } + if eff: + budget_map = {"low": 1024, "medium": 2048, "high": 4096} + budget = budget_map.get(eff, 2048) + gen_config["thinkingConfig"] = {"thinkingBudget": budget} + + parts: List[Dict[str, Any]] = [{"text": prompt}] + if image_path and os.path.isfile(image_path): + try: + mime = "image/jpeg" + lower_p = image_path.lower() + if lower_p.endswith(".png"): + mime = "image/png" + elif lower_p.endswith(".webp"): + mime = "image/webp" + with open(image_path, "rb") as f: + b64 = base64.b64encode(f.read()).decode("utf-8") + parts.append({"inlineData": {"mimeType": mime, "data": b64}}) + except Exception as img_err: + logger.warning(f"Failed to read image {image_path} for Gemini inlineData: {img_err}") - async def _call_gemini(self, prompt: str, system_prompt: Optional[str] = None) -> Dict[str, Any]: - url = f"https://generativelanguage.googleapis.com/v1beta/models/{self.model}:generateContent?key={self.api_key}" payload: Dict[str, Any] = { "contents": [ { - "parts": [{"text": prompt}] + "parts": parts } ], - "generationConfig": { - "responseMimeType": "application/json", - "temperature": 0.2, - } + "generationConfig": gen_config } if system_prompt: payload["systemInstruction"] = { @@ -63,31 +337,138 @@ class LLMClient: resp.raise_for_status() data = resp.json() raw_text = data["candidates"][0]["content"]["parts"][0]["text"] - return json.loads(raw_text) + return self._parse_json_content(raw_text) - async def _call_openai(self, prompt: str, system_prompt: Optional[str] = None) -> Dict[str, Any]: - base = (self.base_url or "https://api.orcarouter.ai/v1").rstrip("/") - url = base if base.endswith("/chat/completions") else f"{base}/chat/completions" + def _resolve_openai_url(self, base_url: Optional[str], is_agy: bool = False) -> str: + base = (base_url or "").strip() + if not base: + if is_agy or self.provider == "agy": + base = "http://host.docker.internal:8088/v1" if os.path.exists("/.dockerenv") else "http://localhost:8088/v1" + else: + base = "https://api.orcarouter.ai/v1" + if os.path.exists("/.dockerenv"): + base = base.replace("localhost", "host.docker.internal").replace("127.0.0.1", "host.docker.internal") + base = base.rstrip("/") + return base if base.endswith("/chat/completions") else f"{base}/chat/completions" + + async def _call_openai( + self, + prompt: str, + system_prompt: Optional[str] = None, + model: Optional[str] = None, + base_url: Optional[str] = None, + api_key: Optional[str] = None, + reasoning_effort: Optional[str] = None, + is_agy: bool = False, + image_path: Optional[str] = None, + ) -> Dict[str, Any]: + model = model or self.model + url = self._resolve_openai_url(base_url or self.base_url, is_agy=is_agy) + eff = (reasoning_effort or self.reasoning_effort or "").strip().lower() + key = (api_key or self.api_key or "").strip() headers = { - "Authorization": f"Bearer {self.api_key}", "Content-Type": "application/json" } + if key: + headers["Authorization"] = f"Bearer {key}" + messages = [] if system_prompt: messages.append({"role": "system", "content": system_prompt}) - messages.append({"role": "user", "content": prompt}) - payload = { - "model": self.model, + # Format user content: text + optional image + if image_path and os.path.isfile(image_path): + try: + mime = "image/jpeg" + lower_p = image_path.lower() + if lower_p.endswith(".png"): + mime = "image/png" + elif lower_p.endswith(".webp"): + mime = "image/webp" + with open(image_path, "rb") as f: + b64 = base64.b64encode(f.read()).decode("utf-8") + user_content: List[Dict[str, Any]] = [ + {"type": "text", "text": prompt}, + {"type": "image_url", "image_url": {"url": f"data:{mime};base64,{b64}"}} + ] + messages.append({"role": "user", "content": user_content}) + except Exception as img_err: + logger.warning(f"Failed to read image {image_path} for vision payload: {img_err}") + messages.append({"role": "user", "content": prompt}) + else: + messages.append({"role": "user", "content": prompt}) + + payload: Dict[str, Any] = { + "model": model, "messages": messages, "response_format": {"type": "json_object"}, "temperature": 0.2, } + if eff in ("low", "medium", "high"): + payload["reasoning_effort"] = eff async with httpx.AsyncClient(timeout=60.0) as client: - resp = await client.post(url, headers=headers, json=payload) - resp.raise_for_status() + try: + resp = await client.post(url, headers=headers, json=payload) + resp.raise_for_status() + except httpx.HTTPStatusError as e: + # Some local servers don't accept response_format: json_object + if e.response.status_code == 400 and "response_format" in payload: + payload.pop("response_format", None) + resp = await client.post(url, headers=headers, json=payload) + resp.raise_for_status() + else: + raise + data = resp.json() content = data["choices"][0]["message"]["content"] - return json.loads(content) + return self._parse_json_content(content) + + async def _call_agy_cli( + self, + prompt: str, + system_prompt: Optional[str] = None, + effort: Optional[str] = None, + image_path: Optional[str] = None, + ) -> Dict[str, Any]: + agy_bin = shutil.which("agy") or ("/usr/local/bin/agy" if os.path.exists("/usr/local/bin/agy") else ("/home/mamad/.local/bin/agy" if os.path.exists("/home/mamad/.local/bin/agy") else None)) + if not agy_bin: + raise RuntimeError("agy binary not found in PATH or mounted paths") + full_prompt = f"System Instruction: {system_prompt}\n\nUser Prompt: {prompt}" if system_prompt else prompt + cmd = [agy_bin, "-p", full_prompt, "--output-format", "json"] + if image_path and os.path.isfile(image_path): + cmd.extend(["-f", image_path]) + eff = (effort or self.reasoning_effort or "").strip().lower() + if eff in ("low", "medium", "high"): + cmd.extend(["--effort", eff]) + env = os.environ.copy() + if not env.get("HOME"): + env["HOME"] = "/root" + proc = await asyncio.create_subprocess_exec( + *cmd, + stdout=asyncio.subprocess.PIPE, + stderr=asyncio.subprocess.PIPE, + env=env + ) + stdout, stderr = await proc.communicate() + if proc.returncode != 0: + raise RuntimeError(f"agy execution failed (code {proc.returncode}): {stderr.decode('utf-8')}") + data = json.loads(stdout.decode("utf-8")) + raw_text = data.get("response", "") + return self._parse_json_content(raw_text) + + async def _call_agy_sdk(self, prompt: str, system_prompt: Optional[str] = None) -> Dict[str, Any]: + if not HAS_AGY_SDK: + raise RuntimeError("google-antigravity SDK is not installed") + config = AgyLocalAgentConfig(system_instructions=system_prompt or "") + async with AgyAgent(config) as agent: + response = await agent.chat(prompt) + chunks = [] + async for token in response: + chunks.append(token) + raw_text = "".join(chunks) + return self._parse_json_content(raw_text) + + + diff --git a/services/ai_processor.py b/services/ai_processor.py index 7ff1403..8d69b32 100644 --- a/services/ai_processor.py +++ b/services/ai_processor.py @@ -1,5 +1,7 @@ +import os import logging import json +from dataclasses import dataclass from typing import List, Optional, Dict, Any from db.models import Post, TargetChannel from db.repository import Repository @@ -9,54 +11,293 @@ from core.error_logger import log_exception logger = logging.getLogger(__name__) -CHANNEL_REWRITE_SYSTEM_PROMPT = """ -You are a professional Persian Telegram copywriter and editor. -Your job is to rewrite the provided raw post specifically for the target channel: "{channel_title}". +@dataclass +class RewriteResult: + decision: str = "accept" # "accept" or "reject" + rejection_reason: str = "" + rewritten_text: str = "" -CHANNEL PERSONALITY & TONE GUIDELINES: -{personality} + @property + def is_rejected(self) -> bool: + return self.decision.lower() == "reject" + + def __str__(self) -> str: + return self.rewritten_text + + +SUPPORTED_LANGUAGES = { + "fa": { + "name": "Persian / Farsi", + "instruction": "Translate or rewrite into natural, highly engaging, and fluent Persian (فارسی روان، جذاب و حرفه‌ای).", + "default_personality": "لحن رسمی، جذاب و روان به همراه ایموجی‌های مرتبط و پاراگراف‌بندی مرتب.", + "label": "فارسی 🇮🇷", + }, + "en": { + "name": "English", + "instruction": "Translate or rewrite into natural, highly engaging, and fluent English.", + "default_personality": "Professional, engaging, and clear tone with appropriate emojis and clean paragraph formatting.", + "label": "English 🇬🇧", + }, + "es": { + "name": "Spanish / Español", + "instruction": "Translate or rewrite into natural, highly engaging, and fluent Spanish (español natural, fluido y profesional).", + "default_personality": "Tono profesional, atractivo y fluido con emojis apropiados y párrafos claros.", + "label": "Español 🇪🇸", + }, + "ar": { + "name": "Arabic / العربية", + "instruction": "Translate or rewrite into natural, highly engaging, and fluent Arabic (عربي فصيح وسلس وجذاب).", + "default_personality": "أسلوب مهني وجذاب وسلس مع إيموجي مناسبة وفقرات واضحة.", + "label": "العربية 🇸🇦", + }, + "tr": { + "name": "Turkish / Türkçe", + "instruction": "Translate or rewrite into natural, highly engaging, and fluent Turkish (akıcı, doğal ve ilgi çekici Türkçe).", + "default_personality": "İlgili emojiler ve düzenli paragraflarla profesyonel, akıcı ve ilgi çekici bir ton.", + "label": "Türkçe 🇹🇷", + }, +} + +# The prompt embeds a literal JSON example, so the placeholders are substituted by name +# rather than through str.format(), which would try to read the braces as fields. +CHANNEL_REWRITE_SYSTEM_PROMPT = """ +You are a professional __LANGUAGE_NAME__ Telegram copywriter and content adapter. +Your job is to transform and rewrite raw posts specifically to match the target channel's personality and tone: "__CHANNEL_TITLE__". + +TARGET CHANNEL DESIRED OUTPUT PERSONALITY & TONE: +__PERSONALITY__ +__CUSTOM_INSTRUCTIONS_BLOCK__ + +IMPORTANT RULES FOR PERSONALITY & TONE: +- You MUST actively, thoroughly rewrite and restyle the post into the requested personality and tone. +- If the personality requires satire, humor, intense sarcasm/slang, formal journalism, or energetic marketing, YOU MUST WRITE IN THAT EXACT VOICE. +- Do NOT output bland, neutral, or unstyled text. +- Do NOT reject incoming posts because their original style, length, or tone is different from this personality. Your task is to transform ANY valid content into this target voice. CRITICAL RULES: 1. Completely REMOVE all original channel usernames (e.g. @source_channel), sponsor tags, author watermarks, and source links. -2. Translate or rewrite into natural, highly engaging, and fluent Persian (فارسی روان، جذاب و حرفه‌ای). -3. Use appropriate emojis and clear paragraph spacing. +2. __LANGUAGE_INSTRUCTION__ +3. Rewrite and format the post to fully embody the target personality with appropriate emojis and clear paragraph spacing. 4. If a custom footer/tag is provided below, append it cleanly at the very end of the post: -{custom_footer} +__CUSTOM_FOOTER__ +5. TELEGRAM CHARACTER LIMIT CONSTRAINT: +__LENGTH_LIMIT_RULE__ +6. POST EVALUATION & REJECTION CRITERIA: +- Default decision is "accept". You should adapt and rewrite normal posts even if their original tone or subject is diverse. +- You should ONLY reject a post if it is: + * Pure spam, unrelated scam/gambling/phishing ads + * Completely empty, corrupt, or meaningless text + * Explicitly forbidden by specific negative constraints in the custom commands above +- If REJECTING: set "decision": "reject" and specify the exact reason in "rejection_reason" (in Persian, e.g. "تبلیغات نامرتبط و اسپم"). +- If ACCEPTING (Default): set "decision": "accept", set "rejection_reason": "", and provide the rewritten output in "rewritten_text". Respond ONLY in valid JSON format: { + "decision": "accept", + "rejection_reason": "", "rewritten_text": "..." } """ + + +def build_rewrite_system_prompt( + channel_title: str, + personality: str, + custom_footer: str, + custom_prompt: str = "", + language: str = "fa", + has_media: bool = False, +) -> str: + lang_info = SUPPORTED_LANGUAGES.get(language, SUPPORTED_LANGUAGES["fa"]) + if has_media: + length_rule = "This message contains media (photo/video/document). The total output (including footer) MUST NOT exceed 1000 characters to strictly fit Telegram's 1024-character caption limit." + else: + length_rule = "This is a text-only message. The total output (including footer) MUST NOT exceed 4000 characters to strictly fit Telegram's 4096-character message limit." + + custom_block = "" + if custom_prompt and custom_prompt.strip(): + custom_block = f"\nTARGET CHANNEL SPECIFIC COMMANDS & FORMATTING RULES:\n{custom_prompt.strip()}\n" + + return ( + CHANNEL_REWRITE_SYSTEM_PROMPT + .replace("__CHANNEL_TITLE__", channel_title) + .replace("__LANGUAGE_NAME__", lang_info["name"]) + .replace("__LANGUAGE_INSTRUCTION__", lang_info["instruction"]) + .replace("__PERSONALITY__", personality) + .replace("__CUSTOM_INSTRUCTIONS_BLOCK__", custom_block) + .replace("__CUSTOM_FOOTER__", custom_footer) + .replace("__LENGTH_LIMIT_RULE__", length_rule) + ) + + +VERIFY_REWRITE_SYSTEM_PROMPT = """ +You are a senior editor and quality-assurance reviewer for Telegram channels. +Your task is to double-check and clean up an AI-rewritten post for the target channel: "__CHANNEL_TITLE__". + +TARGET CHANNEL PERSONALITY & TONE (MUST BE PRESERVED): +__PERSONALITY__ +__CUSTOM_INSTRUCTIONS_BLOCK__ + +TARGET LANGUAGE REQUIREMENT: +The ENTIRE post MUST be written 100% in __LANGUAGE_NAME__ (__LANGUAGE_INSTRUCTION__). +NO mixed languages, NO foreign sentences/phrases accidentally left over from other languages, and NO unwanted translations. + +QUALITY & SAFETY AUDIT CHECKLIST: +1. PRESERVE THE PERSONALITY & HUMOR: Do NOT sanitize, formalize, or flatten the draft! Keep the exact tone, humor, satire, slang, and stylistic persona of the draft intact. +2. Fix Language Leaks: If any foreign words or mixed language leaked into the draft, translate them into __LANGUAGE_NAME__ while strictly preserving the persona and tone. +3. Remove any surviving source channel tags, usernames (@...), sponsor links, or URLs from the original post. +4. Keep the custom footer intact if present: +__CUSTOM_FOOTER__ +5. TELEGRAM CHARACTER LIMIT CONSTRAINT: +__LENGTH_LIMIT_RULE__ + +Respond ONLY in valid JSON format: +{ + "final_text": "..." +} +""" + + +def build_verify_system_prompt( + channel_title: str, + personality: str, + custom_footer: str, + custom_prompt: str = "", + language: str = "fa", + has_media: bool = False, +) -> str: + lang_info = SUPPORTED_LANGUAGES.get(language, SUPPORTED_LANGUAGES["fa"]) + if has_media: + length_rule = "This message contains media (photo/video/document). The total output (including footer) MUST NOT exceed 1000 characters to strictly fit Telegram's 1024-character caption limit." + else: + length_rule = "This is a text-only message. The total output (including footer) MUST NOT exceed 4000 characters to strictly fit Telegram's 4096-character message limit." + + custom_block = "" + if custom_prompt and custom_prompt.strip(): + custom_block = f"\nTARGET CHANNEL SPECIFIC COMMANDS & FORMATTING RULES:\n{custom_prompt.strip()}\n" + + return ( + VERIFY_REWRITE_SYSTEM_PROMPT + .replace("__CHANNEL_TITLE__", channel_title) + .replace("__LANGUAGE_NAME__", lang_info["name"]) + .replace("__LANGUAGE_INSTRUCTION__", lang_info["instruction"]) + .replace("__PERSONALITY__", personality) + .replace("__CUSTOM_INSTRUCTIONS_BLOCK__", custom_block) + .replace("__CUSTOM_FOOTER__", custom_footer) + .replace("__LENGTH_LIMIT_RULE__", length_rule) + ) + + class AIProcessor: - def __init__(self, repo: Repository, llm: Optional[LLMClient] = None): + def __init__(self, repo: Repository, llm: Optional[LLMClient] = None, double_check: Optional[bool] = None): self.repo = repo - self.llm = llm or LLMClient() + self.llm = llm or LLMClient(repo=self.repo) + self.explicit_double_check = double_check + if double_check is not None: + self.double_check = double_check + else: + self.double_check = os.getenv("AI_DOUBLE_CHECK", "false").lower() in ("true", "1", "yes") - async def rewrite_for_target(self, raw_text: str, target: TargetChannel) -> str: - """Rewrite raw text according to a specific target channel's personality and custom footer.""" - if not raw_text: - return "" + async def _is_double_check_enabled(self) -> bool: + if self.explicit_double_check is not None: + return self.explicit_double_check + if self.repo: + db_dc = await self.repo.get_setting("ai_double_check") + if db_dc is not None: + return db_dc.strip().lower() in ("true", "1", "yes") + return self.double_check - personality_text = target.personality.strip() if target.personality else "لحن رسمی، جذاب و روان به همراه ایموجی‌های مرتبط و پاراگراف‌بندی مرتب." + async def _verify_rewrite(self, raw_text: str, draft_text: str, target: TargetChannel, has_media: bool = False) -> str: + """Second-pass quality check to fix language mixing, tags, and length constraints without losing personality.""" + lang = getattr(target, "language", "fa") or "fa" + lang_info = SUPPORTED_LANGUAGES.get(lang, SUPPORTED_LANGUAGES["fa"]) + personality_text = target.personality.strip() if target.personality else lang_info["default_personality"] footer_text = target.custom_footer.strip() if target.custom_footer else (f"@{target.username}" if target.username else "") + custom_prompt_text = getattr(target, "custom_prompt", "") or "" - sys_prompt = CHANNEL_REWRITE_SYSTEM_PROMPT.format( + sys_prompt = build_verify_system_prompt( channel_title=target.title or "کانال تلگرام", personality=personality_text, - custom_footer=footer_text + custom_footer=footer_text, + custom_prompt=custom_prompt_text, + language=lang, + has_media=has_media, ) + + user_prompt = ( + f"متن خام اولیه:\n{raw_text}\n\n" + f"پیش‌نویس اولیه بازنویسی‌شده برای بررسی و اصلاح زبان/فرمت:\n{draft_text}" + ) + try: + res = await self.llm.generate_json( + prompt=user_prompt, + system_prompt=sys_prompt, + action_name="verify_target_post" + ) + final_text = res.get("final_text") + if final_text and final_text.strip(): + return final_text.strip() + except Exception as e: + logger.warning(f"Double-check verification failed, keeping initial draft: {e}") + return draft_text + + async def rewrite_for_target( + self, + raw_text: str, + target: TargetChannel, + has_media: bool = False, + image_path: Optional[str] = None, + ) -> RewriteResult: + """Rewrite raw text according to target channel's language, personality, custom prompt commands, length limits, and optional image.""" + if not raw_text and not image_path: + return RewriteResult(decision="reject", rejection_reason="متن پیام و تصویر هر دو خالی هستند", rewritten_text="") + + lang = getattr(target, "language", "fa") or "fa" + lang_info = SUPPORTED_LANGUAGES.get(lang, SUPPORTED_LANGUAGES["fa"]) + personality_text = target.personality.strip() if target.personality else lang_info["default_personality"] + footer_text = target.custom_footer.strip() if target.custom_footer else (f"@{target.username}" if target.username else "") + custom_prompt_text = getattr(target, "custom_prompt", "") or "" + + sys_prompt = build_rewrite_system_prompt( + channel_title=target.title or "کانال تلگرام", + personality=personality_text, + custom_footer=footer_text, + custom_prompt=custom_prompt_text, + language=lang, + has_media=has_media, + ) + + user_prompt_text = f"متن اصلی پست برای بازنویسی و تبدیل به لحن و استایل کانال مقصد:\n\n{raw_text}" if raw_text else "لطفاً با توجه به تصویر پیوست، یک متن جذاب و مناسب برای کانال بنویسید." + try: res = await self.llm.generate_json( - prompt=f"متن اصلی پست برای بازنویسی:\n\n{raw_text}", + prompt=user_prompt_text, system_prompt=sys_prompt, - action_name="rewrite_target_post" + action_name="rewrite_target_post", + image_path=image_path ) - rewritten = res.get("rewritten_text") + decision = str(res.get("decision", "accept")).strip().lower() + rejection_reason = str(res.get("rejection_reason", "")).strip() + rewritten = res.get("rewritten_text") or res.get("text") or "" + + if decision == "reject": + return RewriteResult( + decision="reject", + rejection_reason=rejection_reason or "رد شده طبق ارزیابی هوش مصنوعی", + rewritten_text=rewritten.strip() if rewritten else "" + ) + if rewritten: - return rewritten.strip() + rewritten = rewritten.strip() + if await self._is_double_check_enabled(): + rewritten = await self._verify_rewrite(raw_text, rewritten, target, has_media=has_media) + return RewriteResult( + decision="accept", + rejection_reason="", + rewritten_text=rewritten + ) except Exception as e: await log_exception("ai_processor.rewrite", e, {"target_id": target.id, "target_title": target.title}) @@ -64,4 +305,5 @@ class AIProcessor: fallback = raw_text if footer_text: fallback = f"{fallback}\n\n{footer_text}" - return fallback + return RewriteResult(decision="accept", rejection_reason="", rewritten_text=fallback) + diff --git a/tests/test_prompts.py b/tests/test_prompts.py new file mode 100644 index 0000000..fd1d80d --- /dev/null +++ b/tests/test_prompts.py @@ -0,0 +1,209 @@ +"""The rewrite prompt embeds a literal JSON example, which str.format() must not choke on.""" +import sys + +sys.path.insert(0, "/app") + +from services.ai_processor import build_rewrite_system_prompt, build_verify_system_prompt, AIProcessor +from core.llm import LLMClient +from db.models import TargetChannel +import asyncio + + +def test_prompt_builds_with_literal_json_braces(): + prompt = build_rewrite_system_prompt( + channel_title="کانال تست", + personality="لحن دوستانه", + custom_footer="@test_channel", + language="fa", + has_media=False, + ) + assert "کانال تست" in prompt + assert "لحن دوستانه" in prompt + assert "@test_channel" in prompt + # The JSON schema example must survive verbatim. + assert '"rewritten_text"' in prompt + assert "{" in prompt and "}" in prompt + assert "4000 characters" in prompt + + +def test_prompt_languages(): + for lang, expected in [ + ("en", "English"), + ("es", "Spanish"), + ("ar", "Arabic"), + ("tr", "Turkish"), + ("fa", "Persian"), + ]: + prompt = build_rewrite_system_prompt("Test", "Tone", "@tag", language=lang) + assert expected in prompt, f"Expected {expected} in prompt for lang {lang}" + + +def test_prompt_media_length_limits(): + media_prompt = build_rewrite_system_prompt("Test", "Tone", "@tag", has_media=True) + assert "1000 characters" in media_prompt + assert "1024-character caption limit" in media_prompt + + text_prompt = build_rewrite_system_prompt("Test", "Tone", "@tag", has_media=False) + assert "4000 characters" in text_prompt + assert "4096-character message limit" in text_prompt + + +def test_verify_prompt(): + v_prompt = build_verify_system_prompt("English Channel", "Humorous Tone", "@eng_footer", language="en", has_media=True) + assert "English Channel" in v_prompt + assert "English" in v_prompt + assert "Humorous Tone" in v_prompt + assert "@eng_footer" in v_prompt + assert "1000 characters" in v_prompt + assert '"final_text"' in v_prompt + + + +def test_llm_client_agy_provider(): + client = LLMClient(provider="agy", base_url="http://127.0.0.1:8000/v1") + assert client.provider == "agy" + assert client.model == "antigravity" + assert client.base_url == "http://127.0.0.1:8000/v1" + + +def test_llm_client_reasoning_effort(): + client = LLMClient(reasoning_effort="high") + assert client.reasoning_effort == "high" + + + +class FakeLLM: + def __init__(self): + self.calls = [] + + async def generate_json(self, prompt: str, system_prompt: str = None, action_name: str = "general", *args, **kwargs): + self.calls.append((prompt, system_prompt, action_name)) + if action_name == "rewrite_target_post": + return {"rewritten_text": "Draft with Spanish and English mixed"} + elif action_name == "verify_target_post": + return {"final_text": "Clean English only verified post"} + return {} + + +async def async_test_ai_processor_double_check(): + fake_llm = FakeLLM() + processor = AIProcessor(repo=None, llm=fake_llm, double_check=True) + target = TargetChannel(id=1, channel_id=123, title="Target", username="tgt", language="en") + + result = await processor.rewrite_for_target("Original Raw Post", target, has_media=False) + assert result.rewritten_text == "Clean English only verified post" + assert str(result) == "Clean English only verified post" + assert result.is_rejected is False + assert len(fake_llm.calls) == 2 + assert fake_llm.calls[0][2] == "rewrite_target_post" + assert fake_llm.calls[1][2] == "verify_target_post" + + + +class FakeRepo: + def __init__(self): + self.settings = {"ai_provider": "gemini", "ai_model": "gemini-2.5-flash", "ai_double_check": "true"} + self.logs = [] + self.active_provider = None + + async def get_setting(self, key, default=None): + return self.settings.get(key, default) + + async def set_setting(self, key, value, description=None): + self.settings[key] = value + + async def record_ai_log(self, action_name, provider, model, prompt, system_prompt=None, response_text=None, duration_sec=0.0, status="success", error_message=None): + self.logs.append({ + "action_name": action_name, + "provider": provider, + "model": model, + "prompt": prompt, + "status": status, + }) + return len(self.logs) + + async def get_active_provider_profile(self): + return self.active_provider + + +async def async_test_dynamic_config_and_logging(): + fake_repo = FakeRepo() + client = LLMClient(provider="openai", model="orcarouter/auto", repo=fake_repo) + await client.sync_config_from_repo() + assert client.provider == "gemini" + assert client.model == "gemini-2.5-flash" + + processor = AIProcessor(repo=fake_repo) + assert await processor._is_double_check_enabled() is True + + +async def async_test_provider_profile_sync(): + from db.models import AIProviderProfile + fake_repo = FakeRepo() + fake_repo.active_provider = AIProviderProfile( + id=1, + name="Local AGY", + provider_type="agy", + model="antigravity-turbo", + base_url="http://localhost:8000/v1", + api_key="", + is_active=True + ) + client = LLMClient(provider="openai", model="orcarouter/auto", repo=fake_repo) + await client.sync_config_from_repo() + assert client.provider == "agy" + assert client.model == "antigravity-turbo" + assert client.base_url == "http://localhost:8000/v1" + + +async def async_test_ai_processor_rejection(): + class RejectLLM: + async def generate_json(self, prompt: str, system_prompt: str = None, action_name: str = "general", *args, **kwargs): + return { + "decision": "reject", + "rejection_reason": "تبلیغات نامرتبط و غیرمجاز", + "rewritten_text": "" + } + + processor = AIProcessor(repo=None, llm=RejectLLM()) + target = TargetChannel(id=1, channel_id=123, title="News Channel", username="news", custom_prompt="فقط خبرهای فناوری") + res = await processor.rewrite_for_target("تبلیغ خرید پکیج آموزشی", target) + assert res.is_rejected is True + assert res.rejection_reason == "تبلیغات نامرتبط و غیرمجاز" + assert str(res) == "" + + +def test_custom_target_prompt_in_system_prompt(): + prompt = build_rewrite_system_prompt( + channel_title="Tech News", + personality="رسمی", + custom_footer="@tech_news", + custom_prompt="همیشه تیتر را با ایموجی 🔴 شروع کن و خلاصه‌ای در ۳ خط بنویس", + language="fa" + ) + assert "TARGET CHANNEL SPECIFIC COMMANDS & FORMATTING RULES:" in prompt + assert "همیشه تیتر را با ایموجی 🔴 شروع کن" in prompt + assert "decision" in prompt + assert "rejection_reason" in prompt + + +if __name__ == "__main__": + test_prompt_builds_with_literal_json_braces() + test_prompt_languages() + test_prompt_media_length_limits() + test_verify_prompt() + test_custom_target_prompt_in_system_prompt() + test_llm_client_agy_provider() + test_llm_client_reasoning_effort() + asyncio.run(async_test_ai_processor_double_check()) + asyncio.run(async_test_ai_processor_rejection()) + asyncio.run(async_test_dynamic_config_and_logging()) + asyncio.run(async_test_provider_profile_sync()) + print("All prompt, custom instructions, rejection, and multi-provider profile unit tests passed!") + + + + + + +