File size: 24,827 Bytes
fe52ef9
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
554
from typing import List, Dict, Any, Optional, Generator, Union
from client_utils import get_client, get_model
from agent_config import AgentConfig
import time
import logging
import json

# Update import for Azure Search
from azure.core.credentials import AzureKeyCredential
from azure.search.documents import SearchClient
# Remove the problematic import
# from azure.search.documents.models import Vector
from embedding_client import EmbeddingClient

# Configure logging
logging.basicConfig(
    level=logging.DEBUG,  # Changed from INFO to DEBUG
    format='%(asctime%s - %(name)s - %(levelname)s - %(message)s'
)
logger = logging.getLogger(__name__)

class BaseAgent:
    """Base class for all agents"""
    def __init__(self, config: AgentConfig):
        self.config = config
        # Get model info from config
        self.model_config = config.get_model_config()
        
        # Extract the engine name from model_config (for client lookup)
        self.engine_name = self.model_config.get("model", "grok-2-latest")
        
        # Get the actual model name for API calls
        self.model_name = get_model(self.engine_name)
        
        # Set temperature from config
        self.temperature = self.model_config.get("temperature", 0.0)
        
        # Get the appropriate client for this engine
        self.client = get_client(self.engine_name)
        
        logger.info(f"Initialized {self.__class__.__name__} with engine {self.engine_name}, model {self.model_name}")

    def process(self, messages: List[Dict[str, Any]]) -> str:
        """Process messages and return a response"""
        logger.debug(f"Processing messages: {json.dumps(messages, indent=2)}")
        # Use system prompt from config if defined
        try:
            system_prompt = self.config.get_system_prompt(self.__class__.__name__.lower())
        except Exception:
            try:
                system_prompt = self.config.get_system_prompt("default")
            except Exception:
                system_prompt = None

        # Only prepend the system prompt if not already present
        if system_prompt and (not messages or messages[0].get("role") != "system"):
            messages = [{"role": "system", "content": system_prompt}] + messages

        completion = self.client.chat.completions.create(
            model=self.model_name,
            messages=messages,
            temperature=self.temperature
        )
        logger.debug(f"Received completion: {completion}")
        return completion

    def process_stream(self, messages: List[Dict[str, Any]]) -> Generator[Any, None, None]:
        """Process messages and stream the response"""
        logger.debug(f"Streaming messages: {json.dumps(messages, indent=2)}")
        # Use system prompt from config if defined
        try:
            system_prompt = self.config.get_system_prompt(self.__class__.__name__.lower())
        except Exception:
            try:
                system_prompt = self.config.get_system_prompt("default")
            except Exception:
                system_prompt = None

        # Only prepend the system prompt if not already present
        if system_prompt and (not messages or messages[0].get("role") != "system"):
            messages = [{"role": "system", "content": system_prompt}] + messages

        completion = self.client.chat.completions.create(
            model=self.model_name,
            messages=messages,
            temperature=self.temperature,
            stream=True
        )
        for chunk in completion:
            logger.debug(f"Streaming chunk: {chunk}")
            return chunk

class DataAgent(BaseAgent):
    """Agent specialized in processing queries with different data sources"""
    def __init__(self, config: AgentConfig):
        super().__init__(config)
        self.tool_name = None  # Track which tool is being executed
        self.data_sources = config.get_data_sources()  # Get all data sources

    def get_extra_body(self) -> Optional[Dict[str, Any]]:
        """Get the extra body configuration for the current search type"""
        # Dynamically match the data source based on the current tool name
        for data_source_name, data_source_config in self.data_sources.items():
            if data_source_name in self.tool_name:
                return {"data_sources": [data_source_config]}
        
        # If no matching data source is found, return None
        return None

    def process(self, messages: List[Dict[str, Any]], tool_name: str = None, **kwargs) -> str:
        """Process messages and return a response with the specified tool"""
        self.tool_name = tool_name
        logger.info(f"Processing data agent for tool {tool_name}")
        extra_body = self.get_extra_body()
        # Pass num_results or other kwargs into extra_body if needed
        if extra_body is not None and kwargs:
            extra_body.update({k: v for k, v in kwargs.items() if k in extra_body})
        completion = self.client.chat.completions.create(
            model=self.model_name,
            messages=messages,
            temperature=self.temperature,
            extra_body=extra_body
        )
        logger.debug(f"Data source completion: {completion}")
        return completion

    def process_stream(self, messages: List[Dict[str, Any]], tool_name: str = None, **kwargs) -> Generator[Any, None, None]:
        """Process messages and stream the response with the specified tool"""
        self.tool_name = tool_name
        logger.info(f"Streaming data agent for tool {tool_name}")
        extra_body = self.get_extra_body()
        # Pass num_results or other kwargs into extra_body if needed
        if extra_body is not None and kwargs:
            extra_body.update({k: v for k, v in kwargs.items() if k in extra_body})
        completion = self.client.chat.completions.create(
            model=self.model_name,
            messages=messages,
            temperature=self.temperature,
            stream=True,
            extra_body=extra_body
        )
        for chunk in completion:
            yield chunk

class SearchAgent(BaseAgent):
    """Agent specialized in querying Azure AI Search and Azure Vector Index"""
    def __init__(self, config: AgentConfig):
        super().__init__(config)
        # Extract Azure-specific configurations from config
        self.azure_index = config.get_azure_index()
        self.search_endpoint = self.azure_index.get("search_endpoint")
        self.search_key = self.azure_index.get("search_key")
        self.search_index_name = self.azure_index.get("search_index_name")
        self.vector_index_name = self.azure_index.get("vector_index_name")
        self.embeddings_deployment = self.azure_index.get("embeddings_deployment")
        
        # Initialize the embedding client
        self.embedding_client = EmbeddingClient(
            azure_endpoint=self.azure_index.get("azure_openai_endpoint"),
            api_key=self.azure_index.get("azure_openai_key"),
            deployment=self.embeddings_deployment
        )
        
        # Initialize search clients
        self.search_client = self._create_search_client(self.search_index_name)
        self.vector_client = self._create_search_client(self.vector_index_name)
        
        logger.info(f"Initialized Search Agent with endpoint {self.azure_index.get('search_endpoint')}")
    
    def _create_search_client(self, index_name: str) -> Optional[SearchClient]:
        """Create a search client for the specified index"""
        if not self.search_endpoint or not self.search_key or not index_name:
            logger.warning(f"Missing configuration for search client: endpoint={bool(self.search_endpoint)}, key={bool(self.search_key)}, index={bool(index_name)}")
            return None
            
        try:
            credential = AzureKeyCredential(self.search_key)
            client = SearchClient(
                endpoint=self.search_endpoint,
                index_name=index_name,
                credential=credential
            )
            logger.info(f"Successfully created search client for index {index_name}")
            return client
        except Exception as e:
            logger.error(f"Failed to create search client for index {index_name}: {e}")
            return None
    
    def _generate_embedding(self, text: str) -> List[float]:
        """Generate embedding for vector search using Azure OpenAI"""
        try:
            # Use the embedding client to generate embeddings
            return self.embedding_client.get_embedding(text)
        except Exception as e:
            logger.error(f"Error generating embedding: {e}")
            # Return a fallback embedding (zeros would be neutral in vector space)
            import numpy as np
            embedding = np.zeros(1536)  # Standard dimension for embeddings
            return embedding.tolist()
    
    def query_index(self, query: str, search_type: str, **kwargs) -> Dict[str, Any]:
        """Query the Azure Search index based on search type"""
        
        search_text = kwargs.get("search_text", query)
        logger.info(f"Executing {search_type} search: {search_text}")
        
        if search_type == "vector":
            if not self.vector_client:
                return {"error": "Vector search client not configured"}
            client = self.vector_client
            #embedding = self._generate_embedding(query)
            try:
                results = list(client.search(
                    search_text=search_text,
                    top=kwargs.get("num_results"),
                    include_total_count=True,
                    filter=kwargs.get("filter"),
                    order_by=kwargs.get("orderby"),
                    select=kwargs.get("select"),
                    vector_queries=[
                        {
                            "text": search_text,
                            "fields": kwargs.get("vector_field", "text_vector"),
                            "k": kwargs.get("num_results"),
                            "kind": "text", 
                        }
                    ],
                ))
                
            except Exception as e:
                results = [f"Vector search failed: {e}"]
                logger.info(results)
        else:
            if not self.search_client:
                return {"error": "Index search client not configured"}
            client = self.search_client

            fields = [
                "Id", "Owner", "Name", "CreatedBy", "LastModifiedBy", "ParentId",
                "CreatedDataTime", "LastModifiedDate",
                "Discriminator", "SourceLocationType", "metadata_storage_content_type",
                "metadata_storage_size", "metadata_storage_last_modified",
                "metadata_storage_content_md5", "metadata_storage_name",
                "metadata_storage_path", "metadata_storage_file_extension",
            ]
            
            try:
                raw_results = list(client.search(
                    search_text=search_text,
                    query_type="full",
                    filter=kwargs.get("filter"),
                    order_by=kwargs.get("orderby"),
                    top=kwargs.get("num_results"),
                    include_total_count=True,
                    select=kwargs.get("select"),
                ))
                # Remove key-value pairs with null values from each result dict
                results = [
                    {k: v for k, v in doc.items() if v is not None}
                    for doc in raw_results
                ]
            except Exception as e:
                results = [f"Index search failed: {e}"]
                logger.info(results)
                
        logger.info(f"Search returned {len(results)} results")
        return results
    
    def format_results(self, results: Dict[str, Any]) -> str:
        """Format the search results into readable text"""
        if "error" in results:
            return f"Error: {results['error']}"
        
        if not results or not results.get("value") or len(results["value"]) == 0:
            return "No results found for your query."
            
        docs = results["value"]
        response_lines = [f"Found {len(docs)} relevant documents:"]
        
        for i, doc in enumerate(docs, 1):  # Limit to top 5 for readability
            title = doc.get("metadata_title", doc.get("title", f"Document {i}"))
            content = doc.get("content", "")
            author = doc.get("metadata_author", "Unknown")
            score = doc.get("@search.score", 0)
            url = doc.get("url", "")
            
            response_lines.append(f"\n### {i}. {title}")
            response_lines.append(f"Author: {author} | Score: {score:.2f}")
            
            if url:
                response_lines.append(f"URL: {url}")
                
            # Get a snippet from content (first 150 chars)
            if content:
                snippet = content
                response_lines.append(f"\nPreview: {snippet}")
                
            response_lines.append("-" * 40)
            
        return "\n".join(response_lines)

    def process(self, messages: List[Dict[str, Any]], **kwargs) -> List[Dict[str, Any]]:
        """Process messages by directly querying Azure Search index"""
        # Extract the query from the last user message
        query = ""
        for message in reversed(messages):
            if message["role"] == "user":
                query = message["content"]
                break

        if not query:
            return "No query provided"

        # Determine search type based on tool name or kwargs
        tool_name = kwargs.get("tool_name", "")
        # Use explicit search_mode if provided, else fallback to tool_name logic
        search_mode = kwargs.get("search_mode")
        if search_mode in ("vector", "index"):
            search_type = search_mode
        elif "vector" in tool_name or "content" in tool_name:
            search_type = "vector"
        else:
            search_type = "index"

        results = self.query_index(query, search_type, **kwargs)
        return results

    def process_stream(self, messages: List[Dict[str, Any]], **kwargs) -> Generator[Any, None, None]:
        for chunk in self.process(messages, **kwargs):
            yield chunk

class OrchestratorAgent(BaseAgent):
    """Agent that coordinates between specialized agents using them as tools"""
    def __init__(self, config: AgentConfig):
        # Initialize the base agent with the given config
        super().__init__(config)
        
        self.tools = []
        self.data_agent = None
        self.search_agent = None
        
        # Set up specialized agents if tools are available
        self.tools = config.get_tools()
        self.has_tools = bool(self.tools)
    
        # Initialize the data agent for data source queries
        self.data_agent = DataAgent(config)
        
        # Initialize the search agent for Azure search queries
        self.search_agent = SearchAgent(config)
        
        logger.debug(f"Available tools: {json.dumps(self.tools, indent=2)}")


    def _execute_tool(self, tool_name: str, tool_args: Dict[str, Any], stream: bool = False, messages: List[Dict[str, Any]] = None) -> Union[str, Generator[Any, None, None]]:
        """Execute a tool and return its result"""
        if not self.has_tools:
            error_msg = "No tools available for this engine"
            logger.error(error_msg)
            raise ValueError(error_msg)

        logger.info(f"Executing tool: {tool_name} with arguments: {json.dumps(tool_args, indent=2)}")

        # Use the full message history, appending the tool call as the last user message
        if messages is None:
            messages = []
        if self.config.get_system_prompt(tool_name):
            messages = [{"role": "system", "content": self.config.get_system_prompt(tool_name)}] + messages
            
        tool_agent_map = self.config.get_tool_agent_map()
        agent_type = tool_agent_map.get(tool_name)

        if agent_type == "search_agent":
            if stream:
                return self.search_agent.process_stream(messages, tool_name=tool_name, **tool_args)
            else:
                return self.search_agent.process(messages, tool_name=tool_name, **tool_args)
        elif agent_type == "data_agent":
            if stream:
                return self.data_agent.process_stream(messages, tool_name=tool_name, **tool_args)
            else:
                return self.data_agent.process(messages, tool_name=tool_name, **tool_args)
        elif agent_type == "base_agent":
            if stream:
                return self.process_stream(messages)
            else:
                return self.process(messages)
        else:
            logger.error(f"Unknown tool name: {tool_name}")
            return f"Error: Unknown tool name: {tool_name}"

    def _handle_tool_calls_recursive(self, messages: List[Dict[str, Any]], max_iterations: int = 10) -> Generator[Any, None, None]:
        """Recursively handle tool calls until we get a response without tools"""
        if max_iterations <= 0 or not self.has_tools:
            logger.warning("Max iterations reached or no tools available, stopping recursive tool calls")
            return

        logger.info("Getting model's response with tool calls")
        completion = self.client.chat.completions.create(
            model=self.model_name,
            messages=messages,
            tools=self.tools,
            tool_choice="auto",
            temperature=self.temperature,
            stream=True
        )

        response_text = ""
        tool_calls = []

        for chunk in completion:
            try:
                if chunk.choices and chunk.choices[0].delta:
                    delta = chunk.choices[0].delta
                    if hasattr(delta, 'content') and delta.content:
                        response_text += delta.content
                        logger.debug(f"Assistant content chunk: {delta.content}")
                    if hasattr(delta, 'tool_calls') and delta.tool_calls:
                        for tool_call in delta.tool_calls:
                            if tool_call.index is not None:
                                while len(tool_calls) <= tool_call.index:
                                    tool_calls.append({})
                                if tool_calls[tool_call.index] == {}:
                                    tool_calls[tool_call.index]["index"] = tool_call.index
                                    tool_calls[tool_call.index]["function"] = {"name": "", "arguments": ""}
                                if tool_call.id:
                                    tool_calls[tool_call.index]["id"] = tool_call.id
                                if tool_call.type:
                                    tool_calls[tool_call.index]["type"] = tool_call.type
                                if tool_call.function.name:
                                    tool_calls[tool_call.index]["function"]["name"] = tool_call.function.name
                                if tool_call.function.arguments:
                                    tool_calls[tool_call.index]["function"]["arguments"] += tool_call.function.arguments
                                    logger.debug(f"Tool call arguments chunk: {tool_call.function.arguments}")
                    yield {
                        "role": "assistant",
                        "content": response_text,
                    }
            except Exception as e:
                logger.error(f"Error processing chunk in handle_tool_calls: {e}")
                continue

        if not tool_calls:
            logger.info("No tool calls made, ending process")
            return
        else:
            logger.info(f"Tool calls detected: {json.dumps(tool_calls, indent=2)}")
            for tool_call in tool_calls:
                yield {
                    "role": "assistant",
                    "content": "",
                    "metadata": {
                        "title": tool_call["function"]["name"],
                        "id": tool_call["id"],
                    }
                }
                
        assistant_message = {
            "role": "assistant",
            "content": response_text,
            "tool_calls": tool_calls,
        }

        # Now execute tools and collect tool messages
        tool_messages = []
        for tool_call in tool_calls:
            tool_name = tool_call["function"]["name"]
            try:
                tool_args = eval(tool_call["function"]["arguments"])
                if not isinstance(tool_args, dict):
                    raise ValueError("Tool arguments must be a dictionary")
                logger.debug(f"Parsed tool arguments for {tool_name}: {json.dumps(tool_args, indent=2)}")
            except (SyntaxError, ValueError) as e:
                logger.error(f"Error parsing tool arguments: {e}")
                continue

            response_stream = self._execute_tool(tool_name, tool_args, stream=True, messages=messages)
            response_list = []
            response_text = ""
            
            tool_message = {
                "role": "tool",
                "tool_call_id": tool_call["id"],
                "content": "",
            }
            # Capitalize tool_name and replace underscores with spaces for title
            title = tool_name.replace("_", " ").title()
            for chunk in response_stream:
                try:
                    if hasattr(chunk, 'choices') and chunk.choices and hasattr(chunk.choices[0], 'delta') and hasattr(chunk.choices[0].delta, 'content'):
                        response_text += chunk.choices[0].delta.content or ""
                    elif isinstance(chunk, dict) or isinstance(chunk, list):
                        response_list.append(chunk)
                        response_text += json.dumps(chunk, indent=2)
                    elif isinstance(chunk, str):
                        response_text += chunk
                        
                    tool_message = {
                            "role": "tool",
                            "tool_call_id": tool_call["id"],
                            "content": response_text
                        }
                    yield {
                        "role": "assistant",
                        "content": json.dumps(response_list[:10], indent=2) if response_list else response_text,
                        "metadata": {
                            "title": title,
                            "id": tool_call["id"],
                            "status": "pending",
                        }
                    }
                    
                except Exception as e:
                    response_text = f"Error processing tool response: {e}"
                    logger.error(f"Error processing tool response chunk: {e}")
                    break
                
            yield {
                "role": "assistant",
                "content": json.dumps(response_list[:10], indent=2) if response_list else response_text if response_text else "No results found.",
                "metadata": {
                    "title": title,
                    "id": tool_call["id"],
                    "log": f"(results: {len(response_list)})" if response_list else "",
                    "status": "done",
                }
            }
            
            tool_messages.append(tool_message)
            logger.info(f"Completed tool execution: {tool_name}")

        messages.append(assistant_message)
        messages.extend(tool_messages)

        yield from self._handle_tool_calls_recursive(messages, max_iterations - 1)

    def process_stream(self, messages: List[Dict[str, Any]]) -> Generator[Any, None, None]:
        """Process messages and stream the response"""
        logger.info("Starting orchestration process")
        logger.debug(f"Input messages: {json.dumps(messages, indent=2)}")
        
        # If no tools are available, fall back to the base agent behavior
        if not self.has_tools:
            logger.info("No tools available, falling back to standard model completion")
            yield from super().process_stream(messages)
            return
        
        # Add the system prompt to the messages only if not already present
        system_prompt = self.config.get_system_prompt("orchestrator")
        if not messages or messages[0].get("role") != "system":
            messages = [{"role": "system", "content": system_prompt}, *messages]
        
        # Start the recursive tool call handling
        yield from self._handle_tool_calls_recursive(messages)