File size: 5,138 Bytes
123d40d
 
03da63b
15f4283
 
 
123d40d
15f4283
 
 
 
 
3e9095a
 
15f4283
123d40d
15f4283
3b29cad
 
 
 
 
4e384c9
3e9095a
3b29cad
3e9095a
4e384c9
7e8db03
 
15f4283
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
123d40d
15f4283
 
 
 
 
 
 
 
 
123d40d
15f4283
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
123d40d
 
03da63b
4e384c9
15f4283
 
123d40d
 
03da63b
3e9095a
 
 
 
15f4283
 
 
 
 
 
 
 
 
 
 
 
 
 
4e384c9
15f4283
 
 
 
 
 
 
 
 
123d40d
15f4283
3e9095a
 
 
7e8db03
 
 
 
15f4283
 
4e384c9
15f4283
 
 
 
 
 
 
 
 
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
import os

from langchain.agents import create_agent
from langchain_core.messages import HumanMessage, SystemMessage
from langchain_groq import ChatGroq


def _make_llm(model: str = "openai/gpt-oss-120b", temperature: float = 0.1) -> ChatGroq:
    return ChatGroq(
        model=model,
        temperature=temperature,
        api_key=os.getenv("GROQ_API_KEY"),
        timeout=60,
        max_retries=2,
    )


# Per-role model config. Defaults restricted to models confirmed present on
# this account (openai/gpt-oss-120b and openai/gpt-oss-20b). Override any
# role via env var once you have verified another id with:
#   curl -s https://api.groq.com/openai/v1/models \
#     -H "Authorization: Bearer $GROQ_API_KEY" | jq '.data[].id'
_ROLE_MODELS = {
    "file": os.getenv("GROQ_FILE_MODEL", "openai/gpt-oss-20b"),
    "media": os.getenv("GROQ_MEDIA_MODEL", "openai/gpt-oss-120b"),
    "normalizer": os.getenv("GROQ_NORMALIZER_MODEL", "openai/gpt-oss-20b"),
}


FILE_ANALYSIS_PROMPT = """You are a document analysis specialist.

You handle questions that reference a task_id, URL, or file path pointing to
a document (PDF, XLSX, CSV, TXT, JSON, MD).

Workflow:
1. If given a task_id (no extension, no slash), call `download_task_file`
   to fetch the file and get its local path.
2. Route by extension:
   - .pdf -> read_pdf
   - .xlsx / .xls -> read_excel
   - .csv -> read_csv
   - anything else textual -> read_text_file
3. Answer the user's question from the extracted content. Quote values
   verbatim when possible. Do not speculate."""


MEDIA_ANALYSIS_PROMPT = """You are a media analysis specialist (image, audio, video).

You handle questions that reference a task_id, URL, or file path pointing to
media, or a YouTube URL.

Workflow:
1. If given a task_id, call `download_task_file` first to get a local path.
2. Route by content type:
   - Image (.jpg/.png/.webp/...) -> analyze_image
   - Audio (.mp3/.wav/.m4a/.flac/...) -> transcribe_audio, then reason on the text
   - Video file or non-YouTube video URL -> analyze_video; if audio matters, also transcribe_audio on the same file
   - YouTube URL -> try `youtube_transcript` first; if captions are missing or
     the question is about visuals, also call `analyze_video`
3. Give ONLY the answer the question asks for. No preamble."""


GAIA_NORMALIZER_PROMPT = """You are a strict answer normalizer for the GAIA benchmark.

You receive a QUESTION and a RAW answer produced by an agent. Return ONLY the
final answer in exactly the format the question requires:

- Numbers: digits only, no thousand separators, no units unless the question
  explicitly asks for units. "3,141" -> "3141". "42 seconds" -> "42" unless
  the question asked for a unit.
- Names / entities: just the name, no articles ("the"), no titles, no
  explanatory prefix like "The answer is".
- Yes/No: exactly "Yes" or "No".
- Lists: comma-separated, in the order the question asks (alphabetical,
  chronological, ...). No trailing "and".
- Dates: match the format the question asks for; if unspecified, use
  YYYY-MM-DD.

Never explain, never apologize, never quote the answer.
Output ONLY the normalized answer, on a single line, nothing else."""


class FileAnalysisAgent:
    def __init__(self):
        from tools import (
            download_task_file,
            read_csv,
            read_excel,
            read_pdf,
            read_text_file,
        )

        self.graph = create_agent(
            _make_llm(model=_ROLE_MODELS["file"]),
            tools=[download_task_file, read_pdf, read_excel, read_csv, read_text_file],
            system_prompt=FILE_ANALYSIS_PROMPT,
        )

    def __call__(self, question: str) -> str:
        result = self.graph.invoke(
            {"messages": [("user", question)]},
            config={"recursion_limit": 8},
        )
        return result["messages"][-1].content


class MediaAgent:
    def __init__(self):
        from tools import (
            analyze_image,
            analyze_video,
            download_task_file,
            transcribe_audio,
            youtube_transcript,
        )

        self.graph = create_agent(
            _make_llm(model=_ROLE_MODELS["media"]),
            tools=[
                download_task_file,
                analyze_image,
                analyze_video,
                transcribe_audio,
                youtube_transcript,
            ],
            system_prompt=MEDIA_ANALYSIS_PROMPT,
        )

    def __call__(self, question: str) -> str:
        result = self.graph.invoke(
            {"messages": [("user", question)]},
            config={"recursion_limit": 8},
        )
        return result["messages"][-1].content


class AnswerNormalizer:
    def __init__(self):
        self.llm = _make_llm(model=_ROLE_MODELS["normalizer"], temperature=0.0)

    def __call__(self, question: str, raw_answer: str) -> str:
        response = self.llm.invoke([
            SystemMessage(content=GAIA_NORMALIZER_PROMPT),
            HumanMessage(
                content=f"QUESTION:\n{question}\n\nRAW ANSWER:\n{raw_answer}"
            ),
        ])
        return response.content.strip()