""" Engine wrapper for AutoGEO — exposes clean APIs for the Streamlit UI. """ import sys import os import time from typing import Optional, List, Dict # Ensure the parent is on sys.path so autogeo imports work _parent = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) if _parent not in sys.path: sys.path.insert(0, _parent) from autogeo.rewriters.core import rewrite_document from autogeo.evaluation.generative_engine import ( generate_answer_gemini, generate_answer_gpt, generate_answer_claude, generate_answer_deepseek, generate_answer_doubao, ) from autogeo.evaluation.metrics.geo_score import ( extract_citations_new, impression_pos_count_simple, impression_word_count_simple, impression_wordpos_count_simple, ) from autogeo.utils.deepseek import call_deepseek from autogeo.utils.doubao import call_doubao from autogeo.utils.gemini import call_gemini from autogeo.utils.openai import call_openai from autogeo.utils.anthropic import call_claude from config_ui import get_missing_key_message def _detect_engine(engine_llm: str) -> str: """Map an engine_llm keyword to a canonical engine name.""" llm = engine_llm.lower() if "doubao" in llm or llm.startswith("ep-"): return "doubao" if "deepseek" in llm: return "deepseek" if "gpt" in llm: return "gpt" if "claude" in llm: return "claude" if "gemini" in llm: return "gemini" return "gemini" # --------------------------------------------------------------------------- # Public API # --------------------------------------------------------------------------- def run_rewrite( document: str, engine_llm: str = "doubao", dataset: str = "E-commerce", ) -> str: """Rewrite a document using the chosen engine.""" engine = _detect_engine(engine_llm) msg = get_missing_key_message(engine) if msg: raise ValueError(msg) return rewrite_document( document=document, dataset=dataset, engine_llm=engine_llm, ) def run_evaluate( query: str, text_list: List[str], target_id: int, engine_llm: str = "doubao", ) -> Dict: """Generate an LLM answer from *text_list* and compute GEO scores. Args: query: User question. text_list: List of 5 source documents (target at position *target_id*). target_id: Index of the document we care about. engine_llm: Model keyword. Returns: {"pos": float, "word": float, "wordpos": float} — GEO scores for target doc. """ engine = _detect_engine(engine_llm) # Call the right generate_answer_* function if engine == "deepseek": answer = generate_answer_deepseek(query, text_list, model_name=engine_llm) elif engine == "doubao": answer = generate_answer_doubao(query, text_list, model_name=engine_llm) elif engine == "gpt": answer = generate_answer_gpt(query, text_list, model_name=engine_llm) elif engine == "claude": answer = generate_answer_claude(query, text_list, model_name=engine_llm) else: answer = generate_answer_gemini(query, text_list, model_name="gemini-2.5-flash-lite") citations = extract_citations_new(answer) return { "pos": impression_pos_count_simple(citations)[target_id], "word": impression_word_count_simple(citations)[target_id], "wordpos": impression_wordpos_count_simple(citations)[target_id], } # --------------------------------------------------------------------------- # Iterative refinement — multi-round feedback loop # --------------------------------------------------------------------------- def _call_llm( user_prompt: str, engine_llm: str, system_prompt: str = "", temperature: float = 0.7, ) -> str: """Route a prompt to the right LLM client.""" engine = _detect_engine(engine_llm) msg = get_missing_key_message(engine) if msg: raise ValueError(msg) if engine == "deepseek": return call_deepseek(user_prompt, model_name=engine_llm, system_prompt=system_prompt, temperature=temperature) elif engine == "doubao": return call_doubao(user_prompt, model_name=engine_llm, system_prompt=system_prompt, temperature=temperature) elif engine == "gemini": return call_gemini(user_prompt, model_name="gemini-2.5-flash-lite", system_prompt=system_prompt, temperature=temperature) elif engine == "gpt": return call_openai(user_prompt, model_name=engine_llm, system_prompt=system_prompt, temperature=temperature) elif engine == "claude": return call_claude(user_prompt, model_name=engine_llm, system_prompt=system_prompt, temperature=temperature) else: # Fallback to Doubao return call_doubao(user_prompt, model_name="ep-20260525112247-ccwhk", system_prompt=system_prompt, temperature=temperature) def refine_document( original_doc: str, previous_rewrite: str, scores: Dict[str, float], query: str, engine_llm: str = "doubao", ) -> str: """Refine a rewrite based on evaluation feedback. Args: original_doc: The original document before any rewriting. previous_rewrite: The most recent rewritten version. scores: GEO scores from evaluating *previous_rewrite*. query: The user query for context. engine_llm: Model keyword. Returns: Improved rewritten document. """ # Build targeted feedback issues = [] if scores["pos"] < 0.3: issues.append("引用位置得分极低,开头部分需要立即抓住AI注意力,把核心信息放在最前面") elif scores["pos"] < 0.5: issues.append("引用位置得分偏低,考虑在开头直接给出结论性信息") if scores["word"] < 0.3: issues.append("引用篇幅严重不足,内容需要大幅扩展,增加具体数据、案例、细节") elif scores["word"] < 0.5: issues.append("引用篇幅偏少,可适当补充更多实质性内容") if scores["wordpos"] < 0.3: issues.append("综合引用得分很低,需要同时改善信息位置和内容密度") elif scores["wordpos"] < 0.5: issues.append("综合引用得分有提升空间") issues_text = "\n".join(f"- {i}" for i in issues) if issues else "各维度表现尚可,但仍有优化空间" system_prompt = ( "你是一位专业的 GEO(Generative Engine Optimization)优化专家。" "你的任务是根据评估反馈改进文档,使其在 AI 搜索引擎生成的答案中获得更高的引用可见性。" ) user_prompt = f"""请根据以下评估反馈,优化改写文档。 【用户查询】 {query} 【原始文档】 {original_doc} 【当前改写版本】 {previous_rewrite} 【GEO 评估分数(满分 1.0)】 - 引用位置得分 (pos): {scores['pos']:.4f} — 衡量文档在AI答案中被引用的早晚 - 引用篇幅得分 (word): {scores['word']:.4f} — 衡量文档内容被引用的篇幅占比 - 综合引用得分 (wordpos): {scores['wordpos']:.4f} — 综合位置和篇幅的加权得分 【需要改善的问题】 {issues_text} 请输出优化后的完整文档,保持中文输出,直接输出改写后的文档内容。""" return _call_llm( user_prompt=user_prompt, engine_llm=engine_llm, system_prompt=system_prompt, temperature=0.7, ) def run_rewrite_iterative( document: str, engine_llm: str, query: str, text_list_original: List[str], target_id: int, max_rounds: int = 3, progress_callback=None, ) -> List[Dict]: """Multi-round rewrite + evaluate with feedback. Round 1 uses the original CMU rewrite_document(). Rounds 2+ use refine_document() with evaluation feedback. Args: document: Original document text. engine_llm: Model keyword. query: User query for context. text_list_original: Full text list (target + distractors) for evaluation. target_id: Index of the target document. max_rounds: Maximum number of refinement rounds (1-5). progress_callback: Optional fn(round, status) for UI updates. Returns: List of {round, text, scores} dicts, one per round. """ results = [] # ---- Round 1: Standard rewrite (CMU original) ---- if progress_callback: progress_callback(1, "rewriting") rewritten = run_rewrite(document=document, engine_llm=engine_llm) if progress_callback: progress_callback(1, "evaluating") text_list_eval = [rewritten] + text_list_original[1:] scores = run_evaluate( query=query, text_list=text_list_eval, target_id=target_id, engine_llm=engine_llm, ) results.append({"round": 1, "text": rewritten, "scores": scores}) # ---- Rounds 2+: Refine with feedback ---- for rnd in range(2, max_rounds + 1): prev = results[-1] # Check if already excellent (all scores >= 0.9) if all(v >= 0.9 for v in prev["scores"].values()): break if progress_callback: progress_callback(rnd, "refining") refined = refine_document( original_doc=document, previous_rewrite=prev["text"], scores=prev["scores"], query=query, engine_llm=engine_llm, ) if progress_callback: progress_callback(rnd, "evaluating") text_list_eval = [refined] + text_list_original[1:] scores = run_evaluate( query=query, text_list=text_list_eval, target_id=target_id, engine_llm=engine_llm, ) results.append({"round": rnd, "text": refined, "scores": scores}) # Early stop if no improvement for 2 consecutive rounds if rnd >= 3: prev_prev = results[-2] prev_cur = results[-1] if prev_cur["scores"]["wordpos"] <= prev_prev["scores"]["wordpos"] * 1.02: break return results def generate_distractors( query: str, n: int = 4, engine_llm: str = "doubao", ) -> List[str]: """Use the selected LLM to generate competing search-result summaries. Args: query: The user query. n: Number of distractors to generate. engine_llm: Model keyword used for generation. Returns: List of *n* distractor strings. """ from openai import OpenAI from dotenv import load_dotenv load_dotenv(os.path.join(_parent, "keys.env")) engine = _detect_engine(engine_llm) msg = get_missing_key_message(engine) if msg: raise ValueError(msg) prompt = ( f"针对以下查询,生成 {n} 个不同的网站搜索结果摘要。" f"每个摘要从不同的角度回答查询,模拟真实搜索引擎中竞争网站的内容。" f"每个摘要 200-400 字,使用中文。\n\n" f"查询:{query}\n\n" f'请严格按照 JSON 数组格式输出,每个元素是一个摘要字符串。示例:["摘要1", "摘要2"]' ) if engine == "deepseek": client = OpenAI( api_key=os.getenv("DEEPSEEK_API_KEY"), base_url="https://api.deepseek.com/v1", ) model = engine_llm elif engine == "doubao": client = OpenAI( api_key=os.getenv("DOUBAO_API_KEY"), base_url="https://ark.cn-beijing.volces.com/api/v3", ) model = "ep-20260525112247-ccwhk" if not engine_llm.startswith("ep-") else engine_llm elif engine == "gpt": client = OpenAI(api_key=os.getenv("OPENAI_API_KEY")) model = engine_llm else: # Fallback: use DeepSeek as default distract generator client = OpenAI( api_key=os.getenv("DEEPSEEK_API_KEY"), base_url="https://api.deepseek.com/v1", ) model = "deepseek-chat" import json as _json resp = client.chat.completions.create( model=model, messages=[{"role": "user", "content": prompt}], temperature=0.8, ) raw = resp.choices[0].message.content.strip() # Strip markdown code fences if present (DeepSeek often wraps JSON in ```json ... ```) if raw.startswith("```"): raw = raw.split("\n", 1)[-1] if "\n" in raw else raw[3:] if raw.endswith("```"): raw = raw[:-3].strip() # Try to parse JSON array; if that fails, fall back to line splitting try: items = _json.loads(raw) if isinstance(items, list): return items[:n] except _json.JSONDecodeError: pass # Crude fallback: split by numbered items or double newlines import re parts = re.split(r"\n\d+[\.\、\)]\s*", raw) parts = [p.strip() for p in parts if len(p.strip()) > 20] return parts[:n] if parts else [raw[:400]]