File size: 13,286 Bytes
5a45328
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
"""

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]]