CoolFace
Apppublic

npmaker/Final_Assignment

sourceHugging Faceupdated 1y agoView on Hugging Face
0likes
google_search.py183 linesDownload Raw Back to root
1from smolagents.tools import Tool2 3class GoogleSearchTool(Tool):4    name = "google_web_search"5    description = """Performs a google web search for your query then returns a string of the top search results."""6    inputs = {7        "query": {"type": "string", "description": "The search query to perform."},8        # "filter_year": {9        #     "type": "integer",10        #     "description": "Optionally restrict results to a certain year",11        #     "nullable": True,12        # },13    }14    output_type = "string"15 16    def __init__(self, provider: str = "google_custom"):17        super().__init__()18        import os19        self.provider = provider20        21        if provider == "serpapi":22            self.organic_key = "organic_results"23            api_key_env_name = "SERPAPI_API_KEY"24            self.api_key = os.getenv(api_key_env_name)25            if self.api_key is None:26                raise ValueError(f"Missing API key. Make sure you have '{api_key_env_name}' in your env variables.")27        elif provider == "serper":28            self.organic_key = "organic"29            api_key_env_name = "SERPER_API_KEY"30            self.api_key = os.getenv(api_key_env_name)31            if self.api_key is None:32                raise ValueError(f"Missing API key. Make sure you have '{api_key_env_name}' in your env variables.")33        elif provider == "google_custom":34            self.api_key = os.getenv("GOOGLE_API_KEY")35            self.cx = os.getenv("GOOGLE_CSE_ID")  # Custom Search Engine ID36            if self.api_key is None:37                raise ValueError("Missing API key. Make sure you have 'GOOGLE_API_KEY' in your env variables.")38            if self.cx is None:39                raise ValueError("Missing Custom Search Engine ID. Make sure you have 'GOOGLE_CSE_ID' in your env variables.")40        else:41            raise ValueError(f"Unsupported provider: {provider}")42 43#    def forward(self, query: str, filter_year: int | None = None) -> str:44    def forward(self, query: str) -> str:45        import requests46        47        if self.provider == "serpapi":48            params = {49                "q": query,50                "api_key": self.api_key,51                "engine": "google",52                "google_domain": "google.com",53            }54            base_url = "https://serpapi.com/search.json"55            56            # if filter_year is not None:57            #     params["tbs"] = f"cdr:1,cd_min:01/01/{filter_year},cd_max:12/31/{filter_year}"58                59            response = requests.get(base_url, params=params)60            if response.status_code == 200:61                results = response.json()62            else:63                raise ValueError(response.json())64                65            organic_key = "organic_results"66                67        elif self.provider == "serper":68            params = {69                "q": query,70                "api_key": self.api_key,71            }72            base_url = "https://google.serper.dev/search"73            74            # if filter_year is not None:75            #     params["tbs"] = f"cdr:1,cd_min:01/01/{filter_year},cd_max:12/31/{filter_year}"76                77            response = requests.get(base_url, params=params)78            if response.status_code == 200:79                results = response.json()80            else:81                raise ValueError(response.json())82                83            organic_key = "organic"84            85        elif self.provider == "google_custom":86            params = {87                "q": query,88                "key": self.api_key,89                "cx": self.cx,90                "num": 10,  # Number of results to return91            }92            base_url = "https://www.googleapis.com/customsearch/v1"93            94            # if filter_year is not None:95            #     # Format for Google Custom Search is different96            #     params["sort"] = f"date:r:{filter_year}:{filter_year}"97                98            response = requests.get(base_url, params=params)99            if response.status_code == 200:100                results = response.json()101            else:102                raise ValueError(f"API Error: {response.status_code} - {response.text}")103                104            # Process Google Custom Search format105            if "items" not in results:106                raise Exception(f"No results found for query: '{query}'. Use a less restrictive query.")107                # if filter_year is not None:108                #     raise Exception(109                #         f"No results found for query: '{query}' with filtering on year={filter_year}. Use a less restrictive query or do not filter on year."110                #     )111                # else:112                #     raise Exception(f"No results found for query: '{query}'. Use a less restrictive query.")113            114            # Reformat Google Custom Search results to match the expected format115            formatted_results = []116            for idx, item in enumerate(results["items"]):117                result = {118                    "title": item.get("title", ""),119                    "link": item.get("link", ""),120                    "snippet": item.get("snippet", ""),121                }122                123                # Extract date if available (might be in different fields depending on the result type)124                if "pagemap" in item and "metatags" in item["pagemap"] and len(item["pagemap"]["metatags"]) > 0:125                    date_fields = ["date", "article:published_time", "og:updated_time", "datePublished"]126                    for field in date_fields:127                        if field in item["pagemap"]["metatags"][0]:128                            result["date"] = item["pagemap"]["metatags"][0][field]129                            break130                131                # Extract source from displayLink132                if "displayLink" in item:133                    result["source"] = item["displayLink"]134                    135                formatted_results.append(result)136            137            # Return early with the reformatted results138            web_snippets = []139            for idx, page in enumerate(formatted_results):140                date_published = ""141                if "date" in page:142                    date_published = "\nDate published: " + page["date"]143                source = ""144                if "source" in page:145                    source = "\nSource: " + page["source"]146                snippet = ""147                if "snippet" in page:148                    snippet = "\n" + page["snippet"]149                redacted_version = f"{idx}. [{page['title']}]({page['link']}){date_published}{source}\n{snippet}"150                web_snippets.append(redacted_version)151            152            return "## Search Results\n" + "\n\n".join(web_snippets)153        154        # Process results for SerpAPI and Serper (original code)155        if organic_key not in results.keys():156            raise Exception(f"No results found for query: '{query}'. Use a less restrictive query.")157            # if filter_year is not None:158            #     raise Exception(159            #         f"No results found for query: '{query}' with filtering on year={filter_year}. Use a less restrictive query or do not filter on year."160            #     )161            # else:162            #     raise Exception(f"No results found for query: '{query}'. Use a less restrictive query.")163                164        if len(results[organic_key]) == 0:165            # year_filter_message = f" with filter year={filter_year}" if filter_year is not None else ""166            # return f"No results found for '{query}'{year_filter_message}. Try with a more general query, or remove the year filter."167            return f"No results found for '{query}'. Try with a more general query."168            169        web_snippets = []170        for idx, page in enumerate(results[organic_key]):171            date_published = ""172            if "date" in page:173                date_published = "\nDate published: " + page["date"]174            source = ""175            if "source" in page:176                source = "\nSource: " + page["source"]177            snippet = ""178            if "snippet" in page:179                snippet = "\n" + page["snippet"]180            redacted_version = f"{idx}. [{page['title']}]({page['link']}){date_published}{source}\n{snippet}"181            web_snippets.append(redacted_version)182            183        return "## Search Results\n" + "\n\n".join(web_snippets)