File size: 12,163 Bytes
bd468ee
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
374834e
bd468ee
 
 
 
374834e
 
 
 
bd468ee
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
374834e
bd468ee
 
374834e
 
 
 
 
 
 
 
 
bd468ee
 
 
374834e
 
 
 
 
 
bd468ee
 
374834e
 
bd468ee
374834e
bd468ee
dfb7db1
 
 
 
 
 
 
 
bd468ee
dfb7db1
 
 
 
374834e
dfb7db1
 
 
 
bd468ee
dfb7db1
 
 
 
bd468ee
 
dfb7db1
 
 
 
374834e
 
bd468ee
dfb7db1
 
 
 
 
 
bd468ee
374834e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
bd468ee
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
374834e
 
 
bd468ee
374834e
 
bd468ee
374834e
 
 
 
 
 
 
 
bd468ee
 
374834e
 
 
 
bd468ee
374834e
 
 
 
bd468ee
374834e
 
 
 
 
 
 
bd468ee
374834e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
bd468ee
 
 
 
 
 
374834e
bd468ee
 
 
 
 
 
 
 
 
 
 
374834e
bd468ee
 
 
 
 
 
 
 
 
 
 
374834e
bd468ee
 
 
 
 
 
 
 
 
 
374834e
bd468ee
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
dfb7db1
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
# Copyright (c) Meta Platforms, Inc. and affiliates.
# All rights reserved.
#
# This source code is licensed under the BSD-style license found in the
# LICENSE file in the root directory of this source tree.

"""

Rag Optimizer Environment Implementation.

The agent acts as a Data Engineer to un-block a broken RAG pipeline.

"""

from copy import deepcopy
from uuid import uuid4
from typing import Dict, Any, List

from openenv.core.env_server.interfaces import Environment
from openenv.core.env_server.types import State

# Import scikit-learn for our Grader (fallback/legacy)
from sklearn.feature_extraction.text import TfidfVectorizer
from sklearn.metrics.pairwise import cosine_similarity
import numpy as np

# Hybrid Search imports
from sentence_transformers import SentenceTransformer
from rank_bm25 import BM25Okapi

try:
    from models import RagOptimizerAction, RagOptimizerObservation
except ImportError:
    from models import RagOptimizerAction, RagOptimizerObservation


class RagOptimizerEnvironment(Environment):
    """

    RAG Optimizer Engine.

    Maintains a simulated Knowledge Base and grades it using TF-IDF.

    """

    SUPPORTS_CONCURRENT_SESSIONS: bool = True

    def __init__(self):
        self._state = State(episode_id=str(uuid4()), step_count=0)
        self.embedding_model = SentenceTransformer('all-MiniLM-L6-v2') 
        self.kb = {}
        self.test_suite = []
        self.dense_vectors = {}
        
        # Load the Real World Datasets
        import json
        import os
        seed_path = os.path.join(os.path.dirname(__file__), "kb_seed.json")
        with open(seed_path, "r") as f:
            self.seed_data = json.load(f)
            
        self._setup_task("easy")
        
    def _setup_task(self, task_id: str):
        if task_id not in ["easy", "medium", "hard"]:
            task_id = "easy"
            
        # Clone the fresh dataset from seed so the agent can destroy it
        self.kb = deepcopy(self.seed_data[task_id])
        
        if task_id == "easy":
            self.test_suite = [
                {"query": "How do I run the server in production with concurrency?", "target_concept": "uvicorn main:app --workers 4"},
                {"query": "Which version of Pydantic does FastAPI use by default?", "target_concept": "defaults to Pydantic v2"}
            ]
            
        elif task_id == "medium":
            self.kb = deepcopy(self.seed_data["hard"])
            self.test_suite = [
                {"query": "Which entity bankruptcy was the climax of the disaster?", "target_concept": "bankruptcy of Lehman Brothers"},
                {"query": "What types of collateralized assets lost their worth?", "target_concept": "Mortgage-backed securities"},
                {"query": "Who did predatory lenders primarily go after?", "target_concept": "low-income homebuyers"}
            ]

        elif task_id == "hard":
            self.kb = {
                # v1 — legacy, wrong routing syntax, poisons retrieval rank for v3
                "fastapi_routing_v1": {
                    "text": "FastAPI routing uses @app.route() decorator similar to Flask. Define routes with methods=['GET','POST']. Pydantic v1 handles all schema validation.",
                    "metadata": {"version": "legacy"}
                },
                # v2 — also legacy, still wrong syntax, further dilutes correct doc
                "fastapi_routing_v2": {
                    "text": "FastAPI routes are defined using @app.route() with type hints added. Validation is done through Pydantic v1 models attached to each endpoint.",
                    "metadata": {"version": "legacy"}
                },
                # v3 — CORRECT current doc, must survive and rank #1 after purge
                "fastapi_routing_v3": {
                    "text": "FastAPI uses @app.get(), @app.post() and other HTTP method decorators for routing. It defaults to Pydantic v2 for request and response validation.",
                    "metadata": {"version": "current"}
                },
            }
            # 15 generic distractors — low semantic overlap, provide ranking noise
            for i in range(15):
                self.kb[f"doc_distractor_{i}"] = {
                    "text": f"General web framework note {i}: always use async handlers for better throughput in high-traffic services.",
                    "metadata": {}
                }
            self.test_suite = [
                # v1 and v2 both contain @app.route() which competes semantically with v3
                # Deleting v1 and v2 pushes v3 to rank 1 → each delete gives a visible reward jump
                {"query": "How are routes defined in FastAPI?",
                 "target_concept": "@app.get(), @app.post() decorators"},
                {"query": "Which Pydantic version does FastAPI v2 use by default?",
                 "target_concept": "defaults to Pydantic v2"},
            ]

        self._rebuild_cache()

    def _rebuild_cache(self):
        """Called whenever KB documents are added, removed, or updated."""
        if not self.kb:
            self.dense_vectors = {}
            return
            
        doc_ids = list(self.kb.keys())
        # Append metadata to text for embedding
        doc_texts = [(self.kb[d]["text"] + " " + " ".join(self.kb[d]["metadata"].values())).strip() for d in doc_ids]
        
        vectors = self.embedding_model.encode(doc_texts, convert_to_tensor=False)
        self.dense_vectors = {doc_id: vectors[i] for i, doc_id in enumerate(doc_ids)}

    def _get_kb_summary(self) -> Dict[str, Dict]:
        """Returns a summary of the KB for the observation."""
        summary = {}
        for k, v in self.kb.items():
            summary[k] = {"metadata": v.get("metadata", {}), "length": len(v.get("text", ""))}
        return summary

    def reset(self, **kwargs) -> RagOptimizerObservation:
        self._state = State(episode_id=str(uuid4()), step_count=0)
        task_id = kwargs.get("task_id") or kwargs.get("task") or "easy"
        self._setup_task(task_id)
        
        return RagOptimizerObservation(
            message=f"RagOptimizerEnv Initialized for task: {task_id}. Resolve conflicts, append metadata, or splinter chunks to win.",
            current_docs=self._get_kb_summary(),
            done=False,
            reward=self._evaluate_kb()
        )

    def _evaluate_kb(self) -> float:
        """The Grader: Evaluates the current KB using Hybrid RRF (BM25 + Semantic MRR)."""
        if not self.kb or not self.test_suite:
            return 0.01
            
        doc_ids = list(self.kb.keys())
        doc_texts = [(self.kb[d]["text"] + " " + " ".join(self.kb[d]["metadata"].values())).strip() for d in doc_ids]
        
        # 1. BM25 Corpus Preparation
        tokenized_corpus = [doc.lower().split() for doc in doc_texts]
        bm25 = BM25Okapi(tokenized_corpus)
        
        # 2. Dense Matrix
        doc_vectors = np.array([self.dense_vectors[d] for d in doc_ids])
        
        mrr_sum = 0.0
        
        for case in self.test_suite:
            # BM25 Search
            tokenized_query = case["query"].lower().split()
            bm25_scores = bm25.get_scores(tokenized_query)
            bm25_ranks = bm25_scores.argsort()[::-1]
            
            # Dense Search
            query_vec = self.embedding_model.encode(case["query"])
            dense_scores = cosine_similarity([query_vec], doc_vectors)[0]
            dense_ranks = dense_scores.argsort()[::-1]
            
            # Reciprocal Rank Fusion (RRF)
            rrf_scores = {d: 0.0 for d in doc_ids}
            k = 60
            for rank, idx in enumerate(bm25_ranks):
                rrf_scores[doc_ids[idx]] += 1.0 / (k + rank + 1)
            for rank, idx in enumerate(dense_ranks):
                rrf_scores[doc_ids[idx]] += 1.0 / (k + rank + 1)
                
            # Grade Top-K fused list using MRR
            ranked_doc_ids = sorted(rrf_scores.keys(), key=lambda x: rrf_scores[x], reverse=True)
            
            req_meta = case.get("required_metadata_key")
            req_meta_val = case.get("required_metadata_value")
            
            case_mrr = 0.0
            for i, doc_id in enumerate(ranked_doc_ids):
                # Is this the true doc?
                doc_text = self.kb[doc_id]["text"].lower()
                if case["target_concept"].lower() in doc_text:
                    valid = True
                    
                    # Medium Task Semantic Check
                    if req_meta and req_meta_val:
                        if self.kb[doc_id]["metadata"].get(req_meta) != req_meta_val:
                            valid = False
                            
                    if valid:
                        case_mrr = 1.0 / (i + 1)  # MRR formula starts at rank 1
                        break
            
            mrr_sum += case_mrr
            
        base_reward = float(mrr_sum / len(self.test_suite))
        
        # Step Cost Penalty calculation (-0.01 per step)
        cost_penalty = self._state.step_count * 0.01
        
        return max(0.01, min(0.99, base_reward - cost_penalty))

    def step(self, action: RagOptimizerAction) -> RagOptimizerObservation:  # type: ignore[override]
        self._state.step_count += 1
        
        msg = ""
        done = False
        reward = 0.01
        
        try:
            if action.action_type == "read_document":
                if action.doc_id in self.kb:
                    msg = f"Content of {action.doc_id}: {self.kb[action.doc_id]['text']}"
                else:
                    msg = f"Error: doc_id {action.doc_id} not found."
            
            elif action.action_type == "delete_document":
                if action.doc_id in self.kb:
                    del self.kb[action.doc_id]
                    self._rebuild_cache()
                    msg = f"Deleted {action.doc_id}."
                else:
                    msg = f"Error: doc_id {action.doc_id} not found."
                    
            elif action.action_type == "update_document":
                if not action.doc_id or not action.text:
                    msg = "Error: doc_id and text required for update_document."
                else:
                    if action.doc_id not in self.kb:
                        self.kb[action.doc_id] = {"text": "", "metadata": {}}
                    self.kb[action.doc_id]["text"] = action.text
                    self._rebuild_cache()
                    msg = f"Updated text for {action.doc_id}."
                    
            elif action.action_type == "add_metadata":
                if not action.doc_id or not action.metadata_key or not action.metadata_value:
                    msg = "Error: doc_id, metadata_key, and metadata_value required."
                else:
                    if action.doc_id not in self.kb:
                        msg = f"Error: doc_id {action.doc_id} not found."
                    else:
                        self.kb[action.doc_id]["metadata"][action.metadata_key] = action.metadata_value
                        self._rebuild_cache()
                        msg = f"Added metadata to {action.doc_id}."
                        
            elif action.action_type == "submit":
                done = True
                reward = self._evaluate_kb()
                msg = f"Evaluation complete. Final reward: {reward:.2f}"
                
        except Exception as e:
            msg = f"Action failed: {str(e)}"

        if not done:
            reward = self._evaluate_kb()

        return RagOptimizerObservation(
            message=msg,
            current_docs=self._get_kb_summary(),
            done=done,
            reward=reward,
        )

    @property
    def state(self) -> State:
        return self._state