refac
This commit is contained in:
@@ -40,20 +40,17 @@ def search_azure(
|
||||
from azure.search.documents import SearchClient
|
||||
except ImportError:
|
||||
log.error(
|
||||
"azure-search-documents package is not installed. "
|
||||
"Install it with: pip install azure-search-documents"
|
||||
'azure-search-documents package is not installed. Install it with: pip install azure-search-documents'
|
||||
)
|
||||
raise ImportError(
|
||||
"azure-search-documents is required for Azure AI Search. "
|
||||
"Install it with: pip install azure-search-documents"
|
||||
'azure-search-documents is required for Azure AI Search. '
|
||||
'Install it with: pip install azure-search-documents'
|
||||
)
|
||||
|
||||
try:
|
||||
# Create search client with API key authentication
|
||||
credential = AzureKeyCredential(api_key)
|
||||
search_client = SearchClient(
|
||||
endpoint=endpoint, index_name=index_name, credential=credential
|
||||
)
|
||||
search_client = SearchClient(endpoint=endpoint, index_name=index_name, credential=credential)
|
||||
|
||||
# Perform the search
|
||||
results = search_client.search(search_text=query, top=count)
|
||||
@@ -68,42 +65,42 @@ def search_azure(
|
||||
|
||||
# Try to find URL field (common names)
|
||||
link = (
|
||||
result_dict.get("url")
|
||||
or result_dict.get("link")
|
||||
or result_dict.get("uri")
|
||||
or result_dict.get("metadata_storage_path")
|
||||
or ""
|
||||
result_dict.get('url')
|
||||
or result_dict.get('link')
|
||||
or result_dict.get('uri')
|
||||
or result_dict.get('metadata_storage_path')
|
||||
or ''
|
||||
)
|
||||
|
||||
# Try to find title field (common names)
|
||||
title = (
|
||||
result_dict.get("title")
|
||||
or result_dict.get("name")
|
||||
or result_dict.get("metadata_title")
|
||||
or result_dict.get("metadata_storage_name")
|
||||
result_dict.get('title')
|
||||
or result_dict.get('name')
|
||||
or result_dict.get('metadata_title')
|
||||
or result_dict.get('metadata_storage_name')
|
||||
or None
|
||||
)
|
||||
|
||||
# Try to find content/snippet field (common names)
|
||||
snippet = (
|
||||
result_dict.get("content")
|
||||
or result_dict.get("snippet")
|
||||
or result_dict.get("description")
|
||||
or result_dict.get("summary")
|
||||
or result_dict.get("text")
|
||||
result_dict.get('content')
|
||||
or result_dict.get('snippet')
|
||||
or result_dict.get('description')
|
||||
or result_dict.get('summary')
|
||||
or result_dict.get('text')
|
||||
or None
|
||||
)
|
||||
|
||||
# Truncate snippet if too long
|
||||
if snippet and len(snippet) > 500:
|
||||
snippet = snippet[:497] + "..."
|
||||
snippet = snippet[:497] + '...'
|
||||
|
||||
if link: # Only add if we found a valid link
|
||||
search_results.append(
|
||||
{
|
||||
"link": link,
|
||||
"title": title,
|
||||
"snippet": snippet,
|
||||
'link': link,
|
||||
'title': title,
|
||||
'snippet': snippet,
|
||||
}
|
||||
)
|
||||
|
||||
@@ -114,13 +111,13 @@ def search_azure(
|
||||
# Convert to SearchResult objects
|
||||
return [
|
||||
SearchResult(
|
||||
link=result["link"],
|
||||
title=result.get("title"),
|
||||
snippet=result.get("snippet"),
|
||||
link=result['link'],
|
||||
title=result.get('title'),
|
||||
snippet=result.get('snippet'),
|
||||
)
|
||||
for result in search_results
|
||||
]
|
||||
|
||||
except Exception as ex:
|
||||
log.error(f"Azure AI Search error: {ex}")
|
||||
log.error(f'Azure AI Search error: {ex}')
|
||||
raise ex
|
||||
|
||||
@@ -21,48 +21,44 @@ def search_bing(
|
||||
filter_list: Optional[list[str]] = None,
|
||||
) -> list[SearchResult]:
|
||||
mkt = locale
|
||||
params = {"q": query, "mkt": mkt, "count": count}
|
||||
headers = {"Ocp-Apim-Subscription-Key": subscription_key}
|
||||
params = {'q': query, 'mkt': mkt, 'count': count}
|
||||
headers = {'Ocp-Apim-Subscription-Key': subscription_key}
|
||||
|
||||
try:
|
||||
response = requests.get(endpoint, headers=headers, params=params)
|
||||
response.raise_for_status()
|
||||
json_response = response.json()
|
||||
results = json_response.get("webPages", {}).get("value", [])
|
||||
results = json_response.get('webPages', {}).get('value', [])
|
||||
if filter_list:
|
||||
results = get_filtered_results(results, filter_list)
|
||||
return [
|
||||
SearchResult(
|
||||
link=result["url"],
|
||||
title=result.get("name"),
|
||||
snippet=result.get("snippet"),
|
||||
link=result['url'],
|
||||
title=result.get('name'),
|
||||
snippet=result.get('snippet'),
|
||||
)
|
||||
for result in results
|
||||
]
|
||||
except Exception as ex:
|
||||
log.error(f"Error: {ex}")
|
||||
log.error(f'Error: {ex}')
|
||||
raise ex
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description="Search Bing from the command line.")
|
||||
parser = argparse.ArgumentParser(description='Search Bing from the command line.')
|
||||
parser.add_argument(
|
||||
"query",
|
||||
'query',
|
||||
type=str,
|
||||
default="Top 10 international news today",
|
||||
help="The search query.",
|
||||
default='Top 10 international news today',
|
||||
help='The search query.',
|
||||
)
|
||||
parser.add_argument('--count', type=int, default=10, help='Number of search results to return.')
|
||||
parser.add_argument('--filter', nargs='*', help='List of filters to apply to the search results.')
|
||||
parser.add_argument(
|
||||
"--count", type=int, default=10, help="Number of search results to return."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--filter", nargs="*", help="List of filters to apply to the search results."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--locale",
|
||||
'--locale',
|
||||
type=str,
|
||||
default="en-US",
|
||||
help="The locale to use for the search, maps to market in api",
|
||||
default='en-US',
|
||||
help='The locale to use for the search, maps to market in api',
|
||||
)
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
@@ -10,43 +10,38 @@ log = logging.getLogger(__name__)
|
||||
|
||||
def _parse_response(response):
|
||||
results = []
|
||||
if "data" in response:
|
||||
data = response["data"]
|
||||
if "webPages" in data:
|
||||
webPages = data["webPages"]
|
||||
if "value" in webPages:
|
||||
if 'data' in response:
|
||||
data = response['data']
|
||||
if 'webPages' in data:
|
||||
webPages = data['webPages']
|
||||
if 'value' in webPages:
|
||||
results = [
|
||||
{
|
||||
"id": item.get("id", ""),
|
||||
"name": item.get("name", ""),
|
||||
"url": item.get("url", ""),
|
||||
"snippet": item.get("snippet", ""),
|
||||
"summary": item.get("summary", ""),
|
||||
"siteName": item.get("siteName", ""),
|
||||
"siteIcon": item.get("siteIcon", ""),
|
||||
"datePublished": item.get("datePublished", "")
|
||||
or item.get("dateLastCrawled", ""),
|
||||
'id': item.get('id', ''),
|
||||
'name': item.get('name', ''),
|
||||
'url': item.get('url', ''),
|
||||
'snippet': item.get('snippet', ''),
|
||||
'summary': item.get('summary', ''),
|
||||
'siteName': item.get('siteName', ''),
|
||||
'siteIcon': item.get('siteIcon', ''),
|
||||
'datePublished': item.get('datePublished', '') or item.get('dateLastCrawled', ''),
|
||||
}
|
||||
for item in webPages["value"]
|
||||
for item in webPages['value']
|
||||
]
|
||||
return results
|
||||
|
||||
|
||||
def search_bocha(
|
||||
api_key: str, query: str, count: int, filter_list: Optional[list[str]] = None
|
||||
) -> list[SearchResult]:
|
||||
def search_bocha(api_key: str, query: str, count: int, filter_list: Optional[list[str]] = None) -> list[SearchResult]:
|
||||
"""Search using Bocha's Search API and return the results as a list of SearchResult objects.
|
||||
|
||||
Args:
|
||||
api_key (str): A Bocha Search API key
|
||||
query (str): The query to search for
|
||||
"""
|
||||
url = "https://api.bochaai.com/v1/web-search?utm_source=ollama"
|
||||
headers = {"Authorization": f"Bearer {api_key}", "Content-Type": "application/json"}
|
||||
url = 'https://api.bochaai.com/v1/web-search?utm_source=ollama'
|
||||
headers = {'Authorization': f'Bearer {api_key}', 'Content-Type': 'application/json'}
|
||||
|
||||
payload = json.dumps(
|
||||
{"query": query, "summary": True, "freshness": "noLimit", "count": count}
|
||||
)
|
||||
payload = json.dumps({'query': query, 'summary': True, 'freshness': 'noLimit', 'count': count})
|
||||
|
||||
response = requests.post(url, headers=headers, data=payload, timeout=5)
|
||||
response.raise_for_status()
|
||||
@@ -56,8 +51,6 @@ def search_bocha(
|
||||
results = get_filtered_results(results, filter_list)
|
||||
|
||||
return [
|
||||
SearchResult(
|
||||
link=result["url"], title=result.get("name"), snippet=result.get("summary")
|
||||
)
|
||||
SearchResult(link=result['url'], title=result.get('name'), snippet=result.get('summary'))
|
||||
for result in results[:count]
|
||||
]
|
||||
|
||||
@@ -8,44 +8,42 @@ from open_webui.retrieval.web.main import SearchResult, get_filtered_results
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def search_brave(
|
||||
api_key: str, query: str, count: int, filter_list: Optional[list[str]] = None
|
||||
) -> list[SearchResult]:
|
||||
def search_brave(api_key: str, query: str, count: int, filter_list: Optional[list[str]] = None) -> list[SearchResult]:
|
||||
"""Search using Brave's Search API and return the results as a list of SearchResult objects.
|
||||
|
||||
Args:
|
||||
api_key (str): A Brave Search API key
|
||||
query (str): The query to search for
|
||||
"""
|
||||
url = "https://api.search.brave.com/res/v1/web/search"
|
||||
url = 'https://api.search.brave.com/res/v1/web/search'
|
||||
headers = {
|
||||
"Accept": "application/json",
|
||||
"Accept-Encoding": "gzip",
|
||||
"X-Subscription-Token": api_key,
|
||||
'Accept': 'application/json',
|
||||
'Accept-Encoding': 'gzip',
|
||||
'X-Subscription-Token': api_key,
|
||||
}
|
||||
params = {"q": query, "count": count}
|
||||
params = {'q': query, 'count': count}
|
||||
|
||||
response = requests.get(url, headers=headers, params=params)
|
||||
|
||||
# Handle 429 rate limiting - Brave free tier allows 1 request/second
|
||||
# If rate limited, wait 1 second and retry once before failing
|
||||
if response.status_code == 429:
|
||||
log.info("Brave Search API rate limited (429), retrying after 1 second...")
|
||||
log.info('Brave Search API rate limited (429), retrying after 1 second...')
|
||||
time.sleep(1)
|
||||
response = requests.get(url, headers=headers, params=params)
|
||||
|
||||
response.raise_for_status()
|
||||
|
||||
json_response = response.json()
|
||||
results = json_response.get("web", {}).get("results", [])
|
||||
results = json_response.get('web', {}).get('results', [])
|
||||
if filter_list:
|
||||
results = get_filtered_results(results, filter_list)
|
||||
|
||||
return [
|
||||
SearchResult(
|
||||
link=result["url"],
|
||||
title=result.get("title"),
|
||||
snippet=result.get("description"),
|
||||
link=result['url'],
|
||||
title=result.get('title'),
|
||||
snippet=result.get('description'),
|
||||
)
|
||||
for result in results[:count]
|
||||
]
|
||||
|
||||
@@ -13,7 +13,7 @@ def search_duckduckgo(
|
||||
count: int,
|
||||
filter_list: Optional[list[str]] = None,
|
||||
concurrent_requests: Optional[int] = None,
|
||||
backend: Optional[str] = "auto",
|
||||
backend: Optional[str] = 'auto',
|
||||
) -> list[SearchResult]:
|
||||
"""
|
||||
Search using DuckDuckGo's Search API and return the results as a list of SearchResult objects.
|
||||
@@ -33,20 +33,18 @@ def search_duckduckgo(
|
||||
|
||||
# Use the ddgs.text() method to perform the search
|
||||
try:
|
||||
search_results = ddgs.text(
|
||||
query, safesearch="moderate", max_results=count, backend=backend
|
||||
)
|
||||
search_results = ddgs.text(query, safesearch='moderate', max_results=count, backend=backend)
|
||||
except RatelimitException as e:
|
||||
log.error(f"RatelimitException: {e}")
|
||||
log.error(f'RatelimitException: {e}')
|
||||
if filter_list:
|
||||
search_results = get_filtered_results(search_results, filter_list)
|
||||
|
||||
# Return the list of search results
|
||||
return [
|
||||
SearchResult(
|
||||
link=result["href"],
|
||||
title=result.get("title"),
|
||||
snippet=result.get("body"),
|
||||
link=result['href'],
|
||||
title=result.get('title'),
|
||||
snippet=result.get('body'),
|
||||
)
|
||||
for result in search_results
|
||||
]
|
||||
|
||||
@@ -7,7 +7,7 @@ from open_webui.retrieval.web.main import SearchResult
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
EXA_API_BASE = "https://api.exa.ai"
|
||||
EXA_API_BASE = 'https://api.exa.ai'
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -31,36 +31,34 @@ def search_exa(
|
||||
count (int): Number of results to return
|
||||
filter_list (Optional[list[str]]): List of domains to filter results by
|
||||
"""
|
||||
log.info(f"Searching with Exa for query: {query}")
|
||||
log.info(f'Searching with Exa for query: {query}')
|
||||
|
||||
headers = {"Authorization": f"Bearer {api_key}", "Content-Type": "application/json"}
|
||||
headers = {'Authorization': f'Bearer {api_key}', 'Content-Type': 'application/json'}
|
||||
|
||||
payload = {
|
||||
"query": query,
|
||||
"numResults": count or 5,
|
||||
"includeDomains": filter_list,
|
||||
"contents": {"text": True, "highlights": True},
|
||||
"type": "auto", # Use the auto search type (keyword or neural)
|
||||
'query': query,
|
||||
'numResults': count or 5,
|
||||
'includeDomains': filter_list,
|
||||
'contents': {'text': True, 'highlights': True},
|
||||
'type': 'auto', # Use the auto search type (keyword or neural)
|
||||
}
|
||||
|
||||
try:
|
||||
response = requests.post(
|
||||
f"{EXA_API_BASE}/search", headers=headers, json=payload
|
||||
)
|
||||
response = requests.post(f'{EXA_API_BASE}/search', headers=headers, json=payload)
|
||||
response.raise_for_status()
|
||||
data = response.json()
|
||||
|
||||
results = []
|
||||
for result in data["results"]:
|
||||
for result in data['results']:
|
||||
results.append(
|
||||
ExaResult(
|
||||
url=result["url"],
|
||||
title=result["title"],
|
||||
text=result["text"],
|
||||
url=result['url'],
|
||||
title=result['title'],
|
||||
text=result['text'],
|
||||
)
|
||||
)
|
||||
|
||||
log.info(f"Found {len(results)} results")
|
||||
log.info(f'Found {len(results)} results')
|
||||
return [
|
||||
SearchResult(
|
||||
link=result.url,
|
||||
@@ -70,5 +68,5 @@ def search_exa(
|
||||
for result in results
|
||||
]
|
||||
except Exception as e:
|
||||
log.error(f"Error searching Exa: {e}")
|
||||
log.error(f'Error searching Exa: {e}')
|
||||
return []
|
||||
|
||||
@@ -24,12 +24,12 @@ def search_external(
|
||||
) -> List[SearchResult]:
|
||||
try:
|
||||
headers = {
|
||||
"User-Agent": "Open WebUI (https://github.com/open-webui/open-webui) RAG Bot",
|
||||
"Authorization": f"Bearer {external_api_key}",
|
||||
'User-Agent': 'Open WebUI (https://github.com/open-webui/open-webui) RAG Bot',
|
||||
'Authorization': f'Bearer {external_api_key}',
|
||||
}
|
||||
headers = include_user_info_headers(headers, user)
|
||||
|
||||
chat_id = getattr(request.state, "chat_id", None)
|
||||
chat_id = getattr(request.state, 'chat_id', None)
|
||||
if chat_id:
|
||||
headers[FORWARD_SESSION_INFO_HEADER_CHAT_ID] = str(chat_id)
|
||||
|
||||
@@ -37,8 +37,8 @@ def search_external(
|
||||
external_url,
|
||||
headers=headers,
|
||||
json={
|
||||
"query": query,
|
||||
"count": count,
|
||||
'query': query,
|
||||
'count': count,
|
||||
},
|
||||
)
|
||||
response.raise_for_status()
|
||||
@@ -47,14 +47,14 @@ def search_external(
|
||||
results = get_filtered_results(results, filter_list)
|
||||
results = [
|
||||
SearchResult(
|
||||
link=result.get("link"),
|
||||
title=result.get("title"),
|
||||
snippet=result.get("snippet"),
|
||||
link=result.get('link'),
|
||||
title=result.get('title'),
|
||||
snippet=result.get('snippet'),
|
||||
)
|
||||
for result in results[:count]
|
||||
]
|
||||
log.info(f"External search results: {results}")
|
||||
log.info(f'External search results: {results}')
|
||||
return results
|
||||
except Exception as e:
|
||||
log.error(f"Error in External search: {e}")
|
||||
log.error(f'Error in External search: {e}')
|
||||
return []
|
||||
|
||||
@@ -17,9 +17,7 @@ def search_firecrawl(
|
||||
from firecrawl import FirecrawlApp
|
||||
|
||||
firecrawl = FirecrawlApp(api_key=firecrawl_api_key, api_url=firecrawl_url)
|
||||
response = firecrawl.search(
|
||||
query=query, limit=count, ignore_invalid_urls=True, timeout=count * 3
|
||||
)
|
||||
response = firecrawl.search(query=query, limit=count, ignore_invalid_urls=True, timeout=count * 3)
|
||||
results = response.web
|
||||
if filter_list:
|
||||
results = get_filtered_results(results, filter_list)
|
||||
@@ -31,8 +29,8 @@ def search_firecrawl(
|
||||
)
|
||||
for result in results[:count]
|
||||
]
|
||||
log.info(f"External search results: {results}")
|
||||
log.info(f'External search results: {results}')
|
||||
return results
|
||||
except Exception as e:
|
||||
log.error(f"Error in External search: {e}")
|
||||
log.error(f'Error in External search: {e}')
|
||||
return []
|
||||
|
||||
@@ -28,11 +28,11 @@ def search_google_pse(
|
||||
Returns:
|
||||
list[SearchResult]: A list of SearchResult objects.
|
||||
"""
|
||||
url = "https://www.googleapis.com/customsearch/v1"
|
||||
url = 'https://www.googleapis.com/customsearch/v1'
|
||||
|
||||
headers = {"Content-Type": "application/json"}
|
||||
headers = {'Content-Type': 'application/json'}
|
||||
if referer:
|
||||
headers["Referer"] = referer
|
||||
headers['Referer'] = referer
|
||||
|
||||
all_results = []
|
||||
start_index = 1 # Google PSE start parameter is 1-based
|
||||
@@ -40,21 +40,19 @@ def search_google_pse(
|
||||
while count > 0:
|
||||
num_results_this_page = min(count, 10) # Google PSE max results per page is 10
|
||||
params = {
|
||||
"cx": search_engine_id,
|
||||
"q": query,
|
||||
"key": api_key,
|
||||
"num": num_results_this_page,
|
||||
"start": start_index,
|
||||
'cx': search_engine_id,
|
||||
'q': query,
|
||||
'key': api_key,
|
||||
'num': num_results_this_page,
|
||||
'start': start_index,
|
||||
}
|
||||
response = requests.request("GET", url, headers=headers, params=params)
|
||||
response = requests.request('GET', url, headers=headers, params=params)
|
||||
response.raise_for_status()
|
||||
json_response = response.json()
|
||||
results = json_response.get("items", [])
|
||||
results = json_response.get('items', [])
|
||||
if results: # check if results are returned. If not, no more pages to fetch.
|
||||
all_results.extend(results)
|
||||
count -= len(
|
||||
results
|
||||
) # Decrement count by the number of results fetched in this page.
|
||||
count -= len(results) # Decrement count by the number of results fetched in this page.
|
||||
start_index += 10 # Increment start index for the next page
|
||||
else:
|
||||
break # No more results from Google PSE, break the loop
|
||||
@@ -64,9 +62,9 @@ def search_google_pse(
|
||||
|
||||
return [
|
||||
SearchResult(
|
||||
link=result["link"],
|
||||
title=result.get("title"),
|
||||
snippet=result.get("snippet"),
|
||||
link=result['link'],
|
||||
title=result.get('title'),
|
||||
snippet=result.get('snippet'),
|
||||
)
|
||||
for result in all_results
|
||||
]
|
||||
|
||||
@@ -7,9 +7,7 @@ from yarl import URL
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def search_jina(
|
||||
api_key: str, query: str, count: int, base_url: str = ""
|
||||
) -> list[SearchResult]:
|
||||
def search_jina(api_key: str, query: str, count: int, base_url: str = '') -> list[SearchResult]:
|
||||
"""
|
||||
Search using Jina's Search API and return the results as a list of SearchResult objects.
|
||||
Args:
|
||||
@@ -21,16 +19,16 @@ def search_jina(
|
||||
Returns:
|
||||
list[SearchResult]: A list of search results
|
||||
"""
|
||||
jina_search_endpoint = base_url if base_url else "https://s.jina.ai/"
|
||||
jina_search_endpoint = base_url if base_url else 'https://s.jina.ai/'
|
||||
|
||||
headers = {
|
||||
"Accept": "application/json",
|
||||
"Content-Type": "application/json",
|
||||
"Authorization": api_key,
|
||||
"X-Retain-Images": "none",
|
||||
'Accept': 'application/json',
|
||||
'Content-Type': 'application/json',
|
||||
'Authorization': api_key,
|
||||
'X-Retain-Images': 'none',
|
||||
}
|
||||
|
||||
payload = {"q": query, "count": count if count <= 10 else 10}
|
||||
payload = {'q': query, 'count': count if count <= 10 else 10}
|
||||
|
||||
url = str(URL(jina_search_endpoint))
|
||||
response = requests.post(url, headers=headers, json=payload)
|
||||
@@ -38,12 +36,12 @@ def search_jina(
|
||||
data = response.json()
|
||||
|
||||
results = []
|
||||
for result in data["data"]:
|
||||
for result in data['data']:
|
||||
results.append(
|
||||
SearchResult(
|
||||
link=result["url"],
|
||||
title=result.get("title"),
|
||||
snippet=result.get("content"),
|
||||
link=result['url'],
|
||||
title=result.get('title'),
|
||||
snippet=result.get('content'),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
@@ -7,9 +7,7 @@ from open_webui.retrieval.web.main import SearchResult, get_filtered_results
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def search_kagi(
|
||||
api_key: str, query: str, count: int, filter_list: Optional[list[str]] = None
|
||||
) -> list[SearchResult]:
|
||||
def search_kagi(api_key: str, query: str, count: int, filter_list: Optional[list[str]] = None) -> list[SearchResult]:
|
||||
"""Search using Kagi's Search API and return the results as a list of SearchResult objects.
|
||||
|
||||
The Search API will inherit the settings in your account, including results personalization and snippet length.
|
||||
@@ -19,23 +17,21 @@ def search_kagi(
|
||||
query (str): The query to search for
|
||||
count (int): The number of results to return
|
||||
"""
|
||||
url = "https://kagi.com/api/v0/search"
|
||||
url = 'https://kagi.com/api/v0/search'
|
||||
headers = {
|
||||
"Authorization": f"Bot {api_key}",
|
||||
'Authorization': f'Bot {api_key}',
|
||||
}
|
||||
params = {"q": query, "limit": count}
|
||||
params = {'q': query, 'limit': count}
|
||||
|
||||
response = requests.get(url, headers=headers, params=params)
|
||||
response.raise_for_status()
|
||||
json_response = response.json()
|
||||
search_results = json_response.get("data", [])
|
||||
search_results = json_response.get('data', [])
|
||||
|
||||
results = [
|
||||
SearchResult(
|
||||
link=result["url"], title=result["title"], snippet=result.get("snippet")
|
||||
)
|
||||
SearchResult(link=result['url'], title=result['title'], snippet=result.get('snippet'))
|
||||
for result in search_results
|
||||
if result["t"] == 0
|
||||
if result['t'] == 0
|
||||
]
|
||||
|
||||
print(results)
|
||||
|
||||
@@ -16,7 +16,7 @@ def get_filtered_results(results, filter_list):
|
||||
filtered_results = []
|
||||
|
||||
for result in results:
|
||||
url = result.get("url") or result.get("link", "") or result.get("href", "")
|
||||
url = result.get('url') or result.get('link', '') or result.get('href', '')
|
||||
if not validators.url(url):
|
||||
continue
|
||||
|
||||
|
||||
@@ -7,32 +7,27 @@ from open_webui.retrieval.web.main import SearchResult, get_filtered_results
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def search_mojeek(
|
||||
api_key: str, query: str, count: int, filter_list: Optional[list[str]] = None
|
||||
) -> list[SearchResult]:
|
||||
def search_mojeek(api_key: str, query: str, count: int, filter_list: Optional[list[str]] = None) -> list[SearchResult]:
|
||||
"""Search using Mojeek's Search API and return the results as a list of SearchResult objects.
|
||||
|
||||
Args:
|
||||
api_key (str): A Mojeek Search API key
|
||||
query (str): The query to search for
|
||||
"""
|
||||
url = "https://api.mojeek.com/search"
|
||||
url = 'https://api.mojeek.com/search'
|
||||
headers = {
|
||||
"Accept": "application/json",
|
||||
'Accept': 'application/json',
|
||||
}
|
||||
params = {"q": query, "api_key": api_key, "fmt": "json", "t": count}
|
||||
params = {'q': query, 'api_key': api_key, 'fmt': 'json', 't': count}
|
||||
|
||||
response = requests.get(url, headers=headers, params=params)
|
||||
response.raise_for_status()
|
||||
json_response = response.json()
|
||||
results = json_response.get("response", {}).get("results", [])
|
||||
results = json_response.get('response', {}).get('results', [])
|
||||
print(results)
|
||||
if filter_list:
|
||||
results = get_filtered_results(results, filter_list)
|
||||
|
||||
return [
|
||||
SearchResult(
|
||||
link=result["url"], title=result.get("title"), snippet=result.get("desc")
|
||||
)
|
||||
for result in results
|
||||
SearchResult(link=result['url'], title=result.get('title'), snippet=result.get('desc')) for result in results
|
||||
]
|
||||
|
||||
@@ -23,30 +23,30 @@ def search_ollama_cloud(
|
||||
count (int): Number of results to return
|
||||
filter_list (Optional[list[str]]): List of domains to filter results by
|
||||
"""
|
||||
log.info(f"Searching with Ollama for query: {query}")
|
||||
log.info(f'Searching with Ollama for query: {query}')
|
||||
|
||||
headers = {"Authorization": f"Bearer {api_key}", "Content-Type": "application/json"}
|
||||
payload = {"query": query, "max_results": count}
|
||||
headers = {'Authorization': f'Bearer {api_key}', 'Content-Type': 'application/json'}
|
||||
payload = {'query': query, 'max_results': count}
|
||||
|
||||
try:
|
||||
response = requests.post(f"{url}/api/web_search", headers=headers, json=payload)
|
||||
response = requests.post(f'{url}/api/web_search', headers=headers, json=payload)
|
||||
response.raise_for_status()
|
||||
data = response.json()
|
||||
|
||||
results = data.get("results", [])
|
||||
log.info(f"Found {len(results)} results")
|
||||
results = data.get('results', [])
|
||||
log.info(f'Found {len(results)} results')
|
||||
|
||||
if filter_list:
|
||||
results = get_filtered_results(results, filter_list)
|
||||
|
||||
return [
|
||||
SearchResult(
|
||||
link=result.get("url", ""),
|
||||
title=result.get("title", ""),
|
||||
snippet=result.get("content", ""),
|
||||
link=result.get('url', ''),
|
||||
title=result.get('title', ''),
|
||||
snippet=result.get('content', ''),
|
||||
)
|
||||
for result in results
|
||||
]
|
||||
except Exception as e:
|
||||
log.error(f"Error searching Ollama: {e}")
|
||||
log.error(f'Error searching Ollama: {e}')
|
||||
return []
|
||||
|
||||
@@ -5,13 +5,13 @@ import requests
|
||||
from open_webui.retrieval.web.main import SearchResult, get_filtered_results
|
||||
|
||||
MODELS = Literal[
|
||||
"sonar",
|
||||
"sonar-pro",
|
||||
"sonar-reasoning",
|
||||
"sonar-reasoning-pro",
|
||||
"sonar-deep-research",
|
||||
'sonar',
|
||||
'sonar-pro',
|
||||
'sonar-reasoning',
|
||||
'sonar-reasoning-pro',
|
||||
'sonar-deep-research',
|
||||
]
|
||||
SEARCH_CONTEXT_USAGE_LEVELS = Literal["low", "medium", "high"]
|
||||
SEARCH_CONTEXT_USAGE_LEVELS = Literal['low', 'medium', 'high']
|
||||
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
@@ -22,8 +22,8 @@ def search_perplexity(
|
||||
query: str,
|
||||
count: int,
|
||||
filter_list: Optional[list[str]] = None,
|
||||
model: MODELS = "sonar",
|
||||
search_context_usage: SEARCH_CONTEXT_USAGE_LEVELS = "medium",
|
||||
model: MODELS = 'sonar',
|
||||
search_context_usage: SEARCH_CONTEXT_USAGE_LEVELS = 'medium',
|
||||
) -> list[SearchResult]:
|
||||
"""Search using Perplexity API and return the results as a list of SearchResult objects.
|
||||
|
||||
@@ -38,66 +38,63 @@ def search_perplexity(
|
||||
"""
|
||||
|
||||
# Handle PersistentConfig object
|
||||
if hasattr(api_key, "__str__"):
|
||||
if hasattr(api_key, '__str__'):
|
||||
api_key = str(api_key)
|
||||
|
||||
try:
|
||||
url = "https://api.perplexity.ai/chat/completions"
|
||||
url = 'https://api.perplexity.ai/chat/completions'
|
||||
|
||||
# Create payload for the API call
|
||||
payload = {
|
||||
"model": model,
|
||||
"messages": [
|
||||
'model': model,
|
||||
'messages': [
|
||||
{
|
||||
"role": "system",
|
||||
"content": "You are a search assistant. Provide factual information with citations.",
|
||||
'role': 'system',
|
||||
'content': 'You are a search assistant. Provide factual information with citations.',
|
||||
},
|
||||
{"role": "user", "content": query},
|
||||
{'role': 'user', 'content': query},
|
||||
],
|
||||
"temperature": 0.2, # Lower temperature for more factual responses
|
||||
"stream": False,
|
||||
"web_search_options": {
|
||||
"search_context_usage": search_context_usage,
|
||||
'temperature': 0.2, # Lower temperature for more factual responses
|
||||
'stream': False,
|
||||
'web_search_options': {
|
||||
'search_context_usage': search_context_usage,
|
||||
},
|
||||
}
|
||||
|
||||
headers = {
|
||||
"Authorization": f"Bearer {api_key}",
|
||||
"Content-Type": "application/json",
|
||||
'Authorization': f'Bearer {api_key}',
|
||||
'Content-Type': 'application/json',
|
||||
}
|
||||
|
||||
# Make the API request
|
||||
response = requests.request("POST", url, json=payload, headers=headers)
|
||||
response = requests.request('POST', url, json=payload, headers=headers)
|
||||
|
||||
# Parse the JSON response
|
||||
json_response = response.json()
|
||||
|
||||
# Extract citations from the response
|
||||
citations = json_response.get("citations", [])
|
||||
citations = json_response.get('citations', [])
|
||||
|
||||
# Create search results from citations
|
||||
results = []
|
||||
for i, citation in enumerate(citations[:count]):
|
||||
# Extract content from the response to use as snippet
|
||||
content = ""
|
||||
if "choices" in json_response and json_response["choices"]:
|
||||
content = ''
|
||||
if 'choices' in json_response and json_response['choices']:
|
||||
if i == 0:
|
||||
content = json_response["choices"][0]["message"]["content"]
|
||||
content = json_response['choices'][0]['message']['content']
|
||||
|
||||
result = {"link": citation, "title": f"Source {i+1}", "snippet": content}
|
||||
result = {'link': citation, 'title': f'Source {i + 1}', 'snippet': content}
|
||||
results.append(result)
|
||||
|
||||
if filter_list:
|
||||
|
||||
results = get_filtered_results(results, filter_list)
|
||||
|
||||
return [
|
||||
SearchResult(
|
||||
link=result["link"], title=result["title"], snippet=result["snippet"]
|
||||
)
|
||||
SearchResult(link=result['link'], title=result['title'], snippet=result['snippet'])
|
||||
for result in results[:count]
|
||||
]
|
||||
|
||||
except Exception as e:
|
||||
log.error(f"Error searching with Perplexity API: {e}")
|
||||
log.error(f'Error searching with Perplexity API: {e}')
|
||||
return []
|
||||
|
||||
@@ -13,7 +13,7 @@ def search_perplexity_search(
|
||||
query: str,
|
||||
count: int,
|
||||
filter_list: Optional[list[str]] = None,
|
||||
api_url: str = "https://api.perplexity.ai/search",
|
||||
api_url: str = 'https://api.perplexity.ai/search',
|
||||
user=None,
|
||||
) -> list[SearchResult]:
|
||||
"""Search using Perplexity API and return the results as a list of SearchResult objects.
|
||||
@@ -29,10 +29,10 @@ def search_perplexity_search(
|
||||
"""
|
||||
|
||||
# Handle PersistentConfig object
|
||||
if hasattr(api_key, "__str__"):
|
||||
if hasattr(api_key, '__str__'):
|
||||
api_key = str(api_key)
|
||||
|
||||
if hasattr(api_url, "__str__"):
|
||||
if hasattr(api_url, '__str__'):
|
||||
api_url = str(api_url)
|
||||
|
||||
try:
|
||||
@@ -40,13 +40,13 @@ def search_perplexity_search(
|
||||
|
||||
# Create payload for the API call
|
||||
payload = {
|
||||
"query": query,
|
||||
"max_results": count,
|
||||
'query': query,
|
||||
'max_results': count,
|
||||
}
|
||||
|
||||
headers = {
|
||||
"Authorization": f"Bearer {api_key}",
|
||||
"Content-Type": "application/json",
|
||||
'Authorization': f'Bearer {api_key}',
|
||||
'Content-Type': 'application/json',
|
||||
}
|
||||
|
||||
# Forward user info headers if user is provided
|
||||
@@ -54,20 +54,17 @@ def search_perplexity_search(
|
||||
headers = include_user_info_headers(headers, user)
|
||||
|
||||
# Make the API request
|
||||
response = requests.request("POST", url, json=payload, headers=headers)
|
||||
response = requests.request('POST', url, json=payload, headers=headers)
|
||||
# Parse the JSON response
|
||||
json_response = response.json()
|
||||
|
||||
# Extract citations from the response
|
||||
results = json_response.get("results", [])
|
||||
results = json_response.get('results', [])
|
||||
|
||||
return [
|
||||
SearchResult(
|
||||
link=result["url"], title=result["title"], snippet=result["snippet"]
|
||||
)
|
||||
for result in results
|
||||
SearchResult(link=result['url'], title=result['title'], snippet=result['snippet']) for result in results
|
||||
]
|
||||
|
||||
except Exception as e:
|
||||
log.error(f"Error searching with Perplexity Search API: {e}")
|
||||
log.error(f'Error searching with Perplexity Search API: {e}')
|
||||
return []
|
||||
|
||||
@@ -21,28 +21,26 @@ def search_searchapi(
|
||||
api_key (str): A searchapi.io API key
|
||||
query (str): The query to search for
|
||||
"""
|
||||
url = "https://www.searchapi.io/api/v1/search"
|
||||
url = 'https://www.searchapi.io/api/v1/search'
|
||||
|
||||
engine = engine or "google"
|
||||
engine = engine or 'google'
|
||||
|
||||
payload = {"engine": engine, "q": query, "api_key": api_key}
|
||||
payload = {'engine': engine, 'q': query, 'api_key': api_key}
|
||||
|
||||
url = f"{url}?{urlencode(payload)}"
|
||||
response = requests.request("GET", url)
|
||||
url = f'{url}?{urlencode(payload)}'
|
||||
response = requests.request('GET', url)
|
||||
|
||||
json_response = response.json()
|
||||
log.info(f"results from searchapi search: {json_response}")
|
||||
log.info(f'results from searchapi search: {json_response}')
|
||||
|
||||
results = sorted(
|
||||
json_response.get("organic_results", []), key=lambda x: x.get("position", 0)
|
||||
)
|
||||
results = sorted(json_response.get('organic_results', []), key=lambda x: x.get('position', 0))
|
||||
if filter_list:
|
||||
results = get_filtered_results(results, filter_list)
|
||||
return [
|
||||
SearchResult(
|
||||
link=result["link"],
|
||||
title=result.get("title"),
|
||||
snippet=result.get("snippet"),
|
||||
link=result['link'],
|
||||
title=result.get('title'),
|
||||
snippet=result.get('snippet'),
|
||||
)
|
||||
for result in results[:count]
|
||||
]
|
||||
|
||||
@@ -38,38 +38,38 @@ def search_searxng(
|
||||
"""
|
||||
|
||||
# Default values for optional parameters are provided as empty strings or None when not specified.
|
||||
language = kwargs.get("language", "all")
|
||||
safesearch = kwargs.get("safesearch", "1")
|
||||
time_range = kwargs.get("time_range", "")
|
||||
categories = "".join(kwargs.get("categories", []))
|
||||
language = kwargs.get('language', 'all')
|
||||
safesearch = kwargs.get('safesearch', '1')
|
||||
time_range = kwargs.get('time_range', '')
|
||||
categories = ''.join(kwargs.get('categories', []))
|
||||
|
||||
params = {
|
||||
"q": query,
|
||||
"format": "json",
|
||||
"pageno": 1,
|
||||
"safesearch": safesearch,
|
||||
"language": language,
|
||||
"time_range": time_range,
|
||||
"categories": categories,
|
||||
"theme": "simple",
|
||||
"image_proxy": 0,
|
||||
'q': query,
|
||||
'format': 'json',
|
||||
'pageno': 1,
|
||||
'safesearch': safesearch,
|
||||
'language': language,
|
||||
'time_range': time_range,
|
||||
'categories': categories,
|
||||
'theme': 'simple',
|
||||
'image_proxy': 0,
|
||||
}
|
||||
|
||||
# Legacy query format
|
||||
if "<query>" in query_url:
|
||||
if '<query>' in query_url:
|
||||
# Strip all query parameters from the URL
|
||||
query_url = query_url.split("?")[0]
|
||||
query_url = query_url.split('?')[0]
|
||||
|
||||
log.debug(f"searching {query_url}")
|
||||
log.debug(f'searching {query_url}')
|
||||
|
||||
response = requests.get(
|
||||
query_url,
|
||||
headers={
|
||||
"User-Agent": "Open WebUI (https://github.com/open-webui/open-webui) RAG Bot",
|
||||
"Accept": "text/html",
|
||||
"Accept-Encoding": "gzip, deflate",
|
||||
"Accept-Language": "en-US,en;q=0.5",
|
||||
"Connection": "keep-alive",
|
||||
'User-Agent': 'Open WebUI (https://github.com/open-webui/open-webui) RAG Bot',
|
||||
'Accept': 'text/html',
|
||||
'Accept-Encoding': 'gzip, deflate',
|
||||
'Accept-Language': 'en-US,en;q=0.5',
|
||||
'Connection': 'keep-alive',
|
||||
},
|
||||
params=params,
|
||||
)
|
||||
@@ -77,13 +77,11 @@ def search_searxng(
|
||||
response.raise_for_status() # Raise an exception for HTTP errors.
|
||||
|
||||
json_response = response.json()
|
||||
results = json_response.get("results", [])
|
||||
sorted_results = sorted(results, key=lambda x: x.get("score", 0), reverse=True)
|
||||
results = json_response.get('results', [])
|
||||
sorted_results = sorted(results, key=lambda x: x.get('score', 0), reverse=True)
|
||||
if filter_list:
|
||||
sorted_results = get_filtered_results(sorted_results, filter_list)
|
||||
return [
|
||||
SearchResult(
|
||||
link=result["url"], title=result.get("title"), snippet=result.get("content")
|
||||
)
|
||||
SearchResult(link=result['url'], title=result.get('title'), snippet=result.get('content'))
|
||||
for result in sorted_results[:count]
|
||||
]
|
||||
|
||||
@@ -21,28 +21,26 @@ def search_serpapi(
|
||||
api_key (str): A serpapi.com API key
|
||||
query (str): The query to search for
|
||||
"""
|
||||
url = "https://serpapi.com/search"
|
||||
url = 'https://serpapi.com/search'
|
||||
|
||||
engine = engine or "google"
|
||||
engine = engine or 'google'
|
||||
|
||||
payload = {"engine": engine, "q": query, "api_key": api_key}
|
||||
payload = {'engine': engine, 'q': query, 'api_key': api_key}
|
||||
|
||||
url = f"{url}?{urlencode(payload)}"
|
||||
response = requests.request("GET", url)
|
||||
url = f'{url}?{urlencode(payload)}'
|
||||
response = requests.request('GET', url)
|
||||
|
||||
json_response = response.json()
|
||||
log.info(f"results from serpapi search: {json_response}")
|
||||
log.info(f'results from serpapi search: {json_response}')
|
||||
|
||||
results = sorted(
|
||||
json_response.get("organic_results", []), key=lambda x: x.get("position", 0)
|
||||
)
|
||||
results = sorted(json_response.get('organic_results', []), key=lambda x: x.get('position', 0))
|
||||
if filter_list:
|
||||
results = get_filtered_results(results, filter_list)
|
||||
return [
|
||||
SearchResult(
|
||||
link=result["link"],
|
||||
title=result.get("title"),
|
||||
snippet=result.get("snippet"),
|
||||
link=result['link'],
|
||||
title=result.get('title'),
|
||||
snippet=result.get('snippet'),
|
||||
)
|
||||
for result in results[:count]
|
||||
]
|
||||
|
||||
@@ -8,34 +8,30 @@ from open_webui.retrieval.web.main import SearchResult, get_filtered_results
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def search_serper(
|
||||
api_key: str, query: str, count: int, filter_list: Optional[list[str]] = None
|
||||
) -> list[SearchResult]:
|
||||
def search_serper(api_key: str, query: str, count: int, filter_list: Optional[list[str]] = None) -> list[SearchResult]:
|
||||
"""Search using serper.dev's API and return the results as a list of SearchResult objects.
|
||||
|
||||
Args:
|
||||
api_key (str): A serper.dev API key
|
||||
query (str): The query to search for
|
||||
"""
|
||||
url = "https://google.serper.dev/search"
|
||||
url = 'https://google.serper.dev/search'
|
||||
|
||||
payload = json.dumps({"q": query})
|
||||
headers = {"X-API-KEY": api_key, "Content-Type": "application/json"}
|
||||
payload = json.dumps({'q': query})
|
||||
headers = {'X-API-KEY': api_key, 'Content-Type': 'application/json'}
|
||||
|
||||
response = requests.request("POST", url, headers=headers, data=payload)
|
||||
response = requests.request('POST', url, headers=headers, data=payload)
|
||||
response.raise_for_status()
|
||||
|
||||
json_response = response.json()
|
||||
results = sorted(
|
||||
json_response.get("organic", []), key=lambda x: x.get("position", 0)
|
||||
)
|
||||
results = sorted(json_response.get('organic', []), key=lambda x: x.get('position', 0))
|
||||
if filter_list:
|
||||
results = get_filtered_results(results, filter_list)
|
||||
return [
|
||||
SearchResult(
|
||||
link=result["link"],
|
||||
title=result.get("title"),
|
||||
snippet=result.get("description"),
|
||||
link=result['link'],
|
||||
title=result.get('title'),
|
||||
snippet=result.get('description'),
|
||||
)
|
||||
for result in results[:count]
|
||||
]
|
||||
|
||||
@@ -12,10 +12,10 @@ def search_serply(
|
||||
api_key: str,
|
||||
query: str,
|
||||
count: int,
|
||||
hl: str = "us",
|
||||
hl: str = 'us',
|
||||
limit: int = 10,
|
||||
device_type: str = "desktop",
|
||||
proxy_location: str = "US",
|
||||
device_type: str = 'desktop',
|
||||
proxy_location: str = 'US',
|
||||
filter_list: Optional[list[str]] = None,
|
||||
) -> list[SearchResult]:
|
||||
"""Search using serper.dev's API and return the results as a list of SearchResult objects.
|
||||
@@ -26,42 +26,40 @@ def search_serply(
|
||||
hl (str): Host Language code to display results in (reference https://developers.google.com/custom-search/docs/xml_results?hl=en#wsInterfaceLanguages)
|
||||
limit (int): The maximum number of results to return [10-100, defaults to 10]
|
||||
"""
|
||||
log.info("Searching with Serply")
|
||||
log.info('Searching with Serply')
|
||||
|
||||
url = "https://api.serply.io/v1/search/"
|
||||
url = 'https://api.serply.io/v1/search/'
|
||||
|
||||
query_payload = {
|
||||
"q": query,
|
||||
"language": "en",
|
||||
"num": limit,
|
||||
"gl": proxy_location.upper(),
|
||||
"hl": hl.lower(),
|
||||
'q': query,
|
||||
'language': 'en',
|
||||
'num': limit,
|
||||
'gl': proxy_location.upper(),
|
||||
'hl': hl.lower(),
|
||||
}
|
||||
|
||||
url = f"{url}{urlencode(query_payload)}"
|
||||
url = f'{url}{urlencode(query_payload)}'
|
||||
headers = {
|
||||
"X-API-KEY": api_key,
|
||||
"X-User-Agent": device_type,
|
||||
"User-Agent": "open-webui",
|
||||
"X-Proxy-Location": proxy_location,
|
||||
'X-API-KEY': api_key,
|
||||
'X-User-Agent': device_type,
|
||||
'User-Agent': 'open-webui',
|
||||
'X-Proxy-Location': proxy_location,
|
||||
}
|
||||
|
||||
response = requests.request("GET", url, headers=headers)
|
||||
response = requests.request('GET', url, headers=headers)
|
||||
response.raise_for_status()
|
||||
|
||||
json_response = response.json()
|
||||
log.info(f"results from serply search: {json_response}")
|
||||
log.info(f'results from serply search: {json_response}')
|
||||
|
||||
results = sorted(
|
||||
json_response.get("results", []), key=lambda x: x.get("realPosition", 0)
|
||||
)
|
||||
results = sorted(json_response.get('results', []), key=lambda x: x.get('realPosition', 0))
|
||||
if filter_list:
|
||||
results = get_filtered_results(results, filter_list)
|
||||
return [
|
||||
SearchResult(
|
||||
link=result["link"],
|
||||
title=result.get("title"),
|
||||
snippet=result.get("description"),
|
||||
link=result['link'],
|
||||
title=result.get('title'),
|
||||
snippet=result.get('description'),
|
||||
)
|
||||
for result in results[:count]
|
||||
]
|
||||
|
||||
@@ -21,26 +21,22 @@ def search_serpstack(
|
||||
query (str): The query to search for
|
||||
https_enabled (bool): Whether to use HTTPS or HTTP for the API request
|
||||
"""
|
||||
url = f"{'https' if https_enabled else 'http'}://api.serpstack.com/search"
|
||||
url = f'{"https" if https_enabled else "http"}://api.serpstack.com/search'
|
||||
|
||||
headers = {"Content-Type": "application/json"}
|
||||
headers = {'Content-Type': 'application/json'}
|
||||
params = {
|
||||
"access_key": api_key,
|
||||
"query": query,
|
||||
'access_key': api_key,
|
||||
'query': query,
|
||||
}
|
||||
|
||||
response = requests.request("POST", url, headers=headers, params=params)
|
||||
response = requests.request('POST', url, headers=headers, params=params)
|
||||
response.raise_for_status()
|
||||
|
||||
json_response = response.json()
|
||||
results = sorted(
|
||||
json_response.get("organic_results", []), key=lambda x: x.get("position", 0)
|
||||
)
|
||||
results = sorted(json_response.get('organic_results', []), key=lambda x: x.get('position', 0))
|
||||
if filter_list:
|
||||
results = get_filtered_results(results, filter_list)
|
||||
return [
|
||||
SearchResult(
|
||||
link=result["url"], title=result.get("title"), snippet=result.get("snippet")
|
||||
)
|
||||
SearchResult(link=result['url'], title=result.get('title'), snippet=result.get('snippet'))
|
||||
for result in results[:count]
|
||||
]
|
||||
|
||||
@@ -26,33 +26,26 @@ def search_sougou(
|
||||
try:
|
||||
cred = credential.Credential(sougou_api_sid, sougou_api_sk)
|
||||
http_profile = HttpProfile()
|
||||
http_profile.endpoint = "tms.tencentcloudapi.com"
|
||||
http_profile.endpoint = 'tms.tencentcloudapi.com'
|
||||
client_profile = ClientProfile()
|
||||
client_profile.http_profile = http_profile
|
||||
params = json.dumps({"Query": query, "Cnt": 20})
|
||||
common_client = CommonClient(
|
||||
"tms", "2020-12-29", cred, "", profile=client_profile
|
||||
)
|
||||
params = json.dumps({'Query': query, 'Cnt': 20})
|
||||
common_client = CommonClient('tms', '2020-12-29', cred, '', profile=client_profile)
|
||||
results = [
|
||||
json.loads(page)
|
||||
for page in common_client.call_json("SearchPro", json.loads(params))[
|
||||
"Response"
|
||||
]["Pages"]
|
||||
json.loads(page) for page in common_client.call_json('SearchPro', json.loads(params))['Response']['Pages']
|
||||
]
|
||||
sorted_results = sorted(
|
||||
results, key=lambda x: x.get("scour", 0.0), reverse=True
|
||||
)
|
||||
sorted_results = sorted(results, key=lambda x: x.get('scour', 0.0), reverse=True)
|
||||
if filter_list:
|
||||
sorted_results = get_filtered_results(sorted_results, filter_list)
|
||||
|
||||
return [
|
||||
SearchResult(
|
||||
link=result.get("url"),
|
||||
title=result.get("title"),
|
||||
snippet=result.get("passage"),
|
||||
link=result.get('url'),
|
||||
title=result.get('title'),
|
||||
snippet=result.get('passage'),
|
||||
)
|
||||
for result in sorted_results[:count]
|
||||
]
|
||||
except TencentCloudSDKException as err:
|
||||
log.error(f"Error in Sougou search: {err}")
|
||||
log.error(f'Error in Sougou search: {err}')
|
||||
return []
|
||||
|
||||
@@ -24,26 +24,26 @@ def search_tavily(
|
||||
Returns:
|
||||
list[SearchResult]: A list of search results
|
||||
"""
|
||||
url = "https://api.tavily.com/search"
|
||||
url = 'https://api.tavily.com/search'
|
||||
headers = {
|
||||
"Content-Type": "application/json",
|
||||
"Authorization": f"Bearer {api_key}",
|
||||
'Content-Type': 'application/json',
|
||||
'Authorization': f'Bearer {api_key}',
|
||||
}
|
||||
data = {"query": query, "max_results": count}
|
||||
data = {'query': query, 'max_results': count}
|
||||
response = requests.post(url, headers=headers, json=data)
|
||||
response.raise_for_status()
|
||||
|
||||
json_response = response.json()
|
||||
|
||||
results = json_response.get("results", [])
|
||||
results = json_response.get('results', [])
|
||||
if filter_list:
|
||||
results = get_filtered_results(results, filter_list)
|
||||
|
||||
return [
|
||||
SearchResult(
|
||||
link=result["url"],
|
||||
title=result.get("title", ""),
|
||||
snippet=result.get("content"),
|
||||
link=result['url'],
|
||||
title=result.get('title', ''),
|
||||
snippet=result.get('content'),
|
||||
)
|
||||
for result in results
|
||||
]
|
||||
|
||||
@@ -67,16 +67,14 @@ def validate_url(url: Union[str, Sequence[str]]):
|
||||
parsed_url = urllib.parse.urlparse(url)
|
||||
|
||||
# Protocol validation - only allow http/https
|
||||
if parsed_url.scheme not in ["http", "https"]:
|
||||
log.warning(
|
||||
f"Blocked non-HTTP(S) protocol: {parsed_url.scheme} in URL: {url}"
|
||||
)
|
||||
if parsed_url.scheme not in ['http', 'https']:
|
||||
log.warning(f'Blocked non-HTTP(S) protocol: {parsed_url.scheme} in URL: {url}')
|
||||
raise ValueError(ERROR_MESSAGES.INVALID_URL)
|
||||
|
||||
# Blocklist check using unified filtering logic
|
||||
if WEB_FETCH_FILTER_LIST:
|
||||
if not is_string_allowed(url, WEB_FETCH_FILTER_LIST):
|
||||
log.warning(f"URL blocked by filter list: {url}")
|
||||
log.warning(f'URL blocked by filter list: {url}')
|
||||
raise ValueError(ERROR_MESSAGES.INVALID_URL)
|
||||
|
||||
if not ENABLE_RAG_LOCAL_WEB_FETCH:
|
||||
@@ -106,29 +104,29 @@ def safe_validate_urls(url: Sequence[str]) -> Sequence[str]:
|
||||
if validate_url(u):
|
||||
valid_urls.append(u)
|
||||
except Exception as e:
|
||||
log.debug(f"Invalid URL {u}: {str(e)}")
|
||||
log.debug(f'Invalid URL {u}: {str(e)}')
|
||||
continue
|
||||
return valid_urls
|
||||
|
||||
|
||||
def extract_metadata(soup, url):
|
||||
metadata = {"source": url}
|
||||
if title := soup.find("title"):
|
||||
metadata["title"] = title.get_text()
|
||||
if description := soup.find("meta", attrs={"name": "description"}):
|
||||
metadata["description"] = description.get("content", "No description found.")
|
||||
if html := soup.find("html"):
|
||||
metadata["language"] = html.get("lang", "No language found.")
|
||||
metadata = {'source': url}
|
||||
if title := soup.find('title'):
|
||||
metadata['title'] = title.get_text()
|
||||
if description := soup.find('meta', attrs={'name': 'description'}):
|
||||
metadata['description'] = description.get('content', 'No description found.')
|
||||
if html := soup.find('html'):
|
||||
metadata['language'] = html.get('lang', 'No language found.')
|
||||
return metadata
|
||||
|
||||
|
||||
def verify_ssl_cert(url: str) -> bool:
|
||||
"""Verify SSL certificate for the given URL."""
|
||||
if not url.startswith("https://"):
|
||||
if not url.startswith('https://'):
|
||||
return True
|
||||
|
||||
try:
|
||||
hostname = url.split("://")[-1].split("/")[0]
|
||||
hostname = url.split('://')[-1].split('/')[0]
|
||||
context = ssl.create_default_context(cafile=certifi.where())
|
||||
with context.wrap_socket(ssl.socket(), server_hostname=hostname) as s:
|
||||
s.connect((hostname, 443))
|
||||
@@ -136,7 +134,7 @@ def verify_ssl_cert(url: str) -> bool:
|
||||
except ssl.SSLError:
|
||||
return False
|
||||
except Exception as e:
|
||||
log.warning(f"SSL verification failed for {url}: {str(e)}")
|
||||
log.warning(f'SSL verification failed for {url}: {str(e)}')
|
||||
return False
|
||||
|
||||
|
||||
@@ -168,14 +166,14 @@ class URLProcessingMixin:
|
||||
async def _safe_process_url(self, url: str) -> bool:
|
||||
"""Perform safety checks before processing a URL."""
|
||||
if self.verify_ssl and not await self._verify_ssl_cert(url):
|
||||
raise ValueError(f"SSL certificate verification failed for {url}")
|
||||
raise ValueError(f'SSL certificate verification failed for {url}')
|
||||
await self._wait_for_rate_limit()
|
||||
return True
|
||||
|
||||
def _safe_process_url_sync(self, url: str) -> bool:
|
||||
"""Synchronous version of safety checks."""
|
||||
if self.verify_ssl and not verify_ssl_cert(url):
|
||||
raise ValueError(f"SSL certificate verification failed for {url}")
|
||||
raise ValueError(f'SSL certificate verification failed for {url}')
|
||||
self._sync_wait_for_rate_limit()
|
||||
return True
|
||||
|
||||
@@ -191,7 +189,7 @@ class SafeFireCrawlLoader(BaseLoader, RateLimitMixin, URLProcessingMixin):
|
||||
api_key: Optional[str] = None,
|
||||
api_url: Optional[str] = None,
|
||||
timeout: Optional[int] = None,
|
||||
mode: Literal["crawl", "scrape", "map"] = "scrape",
|
||||
mode: Literal['crawl', 'scrape', 'map'] = 'scrape',
|
||||
proxy: Optional[Dict[str, str]] = None,
|
||||
params: Optional[Dict] = None,
|
||||
):
|
||||
@@ -216,15 +214,15 @@ class SafeFireCrawlLoader(BaseLoader, RateLimitMixin, URLProcessingMixin):
|
||||
params: The parameters to pass to the Firecrawl API.
|
||||
For more details, visit: https://docs.firecrawl.dev/sdks/python#batch-scrape
|
||||
"""
|
||||
proxy_server = proxy.get("server") if proxy else None
|
||||
proxy_server = proxy.get('server') if proxy else None
|
||||
if trust_env and not proxy_server:
|
||||
env_proxies = urllib.request.getproxies()
|
||||
env_proxy_server = env_proxies.get("https") or env_proxies.get("http")
|
||||
env_proxy_server = env_proxies.get('https') or env_proxies.get('http')
|
||||
if env_proxy_server:
|
||||
if proxy:
|
||||
proxy["server"] = env_proxy_server
|
||||
proxy['server'] = env_proxy_server
|
||||
else:
|
||||
proxy = {"server": env_proxy_server}
|
||||
proxy = {'server': env_proxy_server}
|
||||
self.web_paths = web_paths
|
||||
self.verify_ssl = verify_ssl
|
||||
self.requests_per_second = requests_per_second
|
||||
@@ -240,7 +238,7 @@ class SafeFireCrawlLoader(BaseLoader, RateLimitMixin, URLProcessingMixin):
|
||||
def lazy_load(self) -> Iterator[Document]:
|
||||
"""Load documents using FireCrawl batch_scrape."""
|
||||
log.debug(
|
||||
"Starting FireCrawl batch scrape for %d URLs, mode: %s, params: %s",
|
||||
'Starting FireCrawl batch scrape for %d URLs, mode: %s, params: %s',
|
||||
len(self.web_paths),
|
||||
self.mode,
|
||||
self.params,
|
||||
@@ -251,7 +249,7 @@ class SafeFireCrawlLoader(BaseLoader, RateLimitMixin, URLProcessingMixin):
|
||||
firecrawl = FirecrawlApp(api_key=self.api_key, api_url=self.api_url)
|
||||
result = firecrawl.batch_scrape(
|
||||
self.web_paths,
|
||||
formats=["markdown"],
|
||||
formats=['markdown'],
|
||||
skip_tls_verification=not self.verify_ssl,
|
||||
ignore_invalid_urls=True,
|
||||
remove_base64_images=True,
|
||||
@@ -260,28 +258,26 @@ class SafeFireCrawlLoader(BaseLoader, RateLimitMixin, URLProcessingMixin):
|
||||
**self.params,
|
||||
)
|
||||
|
||||
if result.status != "completed":
|
||||
raise RuntimeError(
|
||||
f"FireCrawl batch scrape did not complete successfully. result: {result}"
|
||||
)
|
||||
if result.status != 'completed':
|
||||
raise RuntimeError(f'FireCrawl batch scrape did not complete successfully. result: {result}')
|
||||
|
||||
for data in result.data:
|
||||
metadata = data.metadata or {}
|
||||
yield Document(
|
||||
page_content=data.markdown or "",
|
||||
metadata={"source": metadata.url or metadata.source_url or ""},
|
||||
page_content=data.markdown or '',
|
||||
metadata={'source': metadata.url or metadata.source_url or ''},
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
if self.continue_on_failure:
|
||||
log.exception(f"Error extracting content from URLs: {e}")
|
||||
log.exception(f'Error extracting content from URLs: {e}')
|
||||
else:
|
||||
raise e
|
||||
|
||||
async def alazy_load(self):
|
||||
"""Async version of lazy_load."""
|
||||
log.debug(
|
||||
"Starting FireCrawl batch scrape for %d URLs, mode: %s, params: %s",
|
||||
'Starting FireCrawl batch scrape for %d URLs, mode: %s, params: %s',
|
||||
len(self.web_paths),
|
||||
self.mode,
|
||||
self.params,
|
||||
@@ -292,7 +288,7 @@ class SafeFireCrawlLoader(BaseLoader, RateLimitMixin, URLProcessingMixin):
|
||||
firecrawl = FirecrawlApp(api_key=self.api_key, api_url=self.api_url)
|
||||
result = firecrawl.batch_scrape(
|
||||
self.web_paths,
|
||||
formats=["markdown"],
|
||||
formats=['markdown'],
|
||||
skip_tls_verification=not self.verify_ssl,
|
||||
ignore_invalid_urls=True,
|
||||
remove_base64_images=True,
|
||||
@@ -301,21 +297,19 @@ class SafeFireCrawlLoader(BaseLoader, RateLimitMixin, URLProcessingMixin):
|
||||
**self.params,
|
||||
)
|
||||
|
||||
if result.status != "completed":
|
||||
raise RuntimeError(
|
||||
f"FireCrawl batch scrape did not complete successfully. result: {result}"
|
||||
)
|
||||
if result.status != 'completed':
|
||||
raise RuntimeError(f'FireCrawl batch scrape did not complete successfully. result: {result}')
|
||||
|
||||
for data in result.data:
|
||||
metadata = data.metadata or {}
|
||||
yield Document(
|
||||
page_content=data.markdown or "",
|
||||
metadata={"source": metadata.url or metadata.source_url or ""},
|
||||
page_content=data.markdown or '',
|
||||
metadata={'source': metadata.url or metadata.source_url or ''},
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
if self.continue_on_failure:
|
||||
log.exception(f"Error extracting content from URLs: {e}")
|
||||
log.exception(f'Error extracting content from URLs: {e}')
|
||||
else:
|
||||
raise e
|
||||
|
||||
@@ -325,7 +319,7 @@ class SafeTavilyLoader(BaseLoader, RateLimitMixin, URLProcessingMixin):
|
||||
self,
|
||||
web_paths: Union[str, List[str]],
|
||||
api_key: str,
|
||||
extract_depth: Literal["basic", "advanced"] = "basic",
|
||||
extract_depth: Literal['basic', 'advanced'] = 'basic',
|
||||
continue_on_failure: bool = True,
|
||||
requests_per_second: Optional[float] = None,
|
||||
verify_ssl: bool = True,
|
||||
@@ -345,15 +339,15 @@ class SafeTavilyLoader(BaseLoader, RateLimitMixin, URLProcessingMixin):
|
||||
proxy: Optional proxy configuration.
|
||||
"""
|
||||
# Initialize proxy configuration if using environment variables
|
||||
proxy_server = proxy.get("server") if proxy else None
|
||||
proxy_server = proxy.get('server') if proxy else None
|
||||
if trust_env and not proxy_server:
|
||||
env_proxies = urllib.request.getproxies()
|
||||
env_proxy_server = env_proxies.get("https") or env_proxies.get("http")
|
||||
env_proxy_server = env_proxies.get('https') or env_proxies.get('http')
|
||||
if env_proxy_server:
|
||||
if proxy:
|
||||
proxy["server"] = env_proxy_server
|
||||
proxy['server'] = env_proxy_server
|
||||
else:
|
||||
proxy = {"server": env_proxy_server}
|
||||
proxy = {'server': env_proxy_server}
|
||||
|
||||
# Store parameters for creating TavilyLoader instances
|
||||
self.web_paths = web_paths if isinstance(web_paths, list) else [web_paths]
|
||||
@@ -376,14 +370,14 @@ class SafeTavilyLoader(BaseLoader, RateLimitMixin, URLProcessingMixin):
|
||||
self._safe_process_url_sync(url)
|
||||
valid_urls.append(url)
|
||||
except Exception as e:
|
||||
log.warning(f"SSL verification failed for {url}: {str(e)}")
|
||||
log.warning(f'SSL verification failed for {url}: {str(e)}')
|
||||
if not self.continue_on_failure:
|
||||
raise e
|
||||
if not valid_urls:
|
||||
if self.continue_on_failure:
|
||||
log.warning("No valid URLs to process after SSL verification")
|
||||
log.warning('No valid URLs to process after SSL verification')
|
||||
return
|
||||
raise ValueError("No valid URLs to process after SSL verification")
|
||||
raise ValueError('No valid URLs to process after SSL verification')
|
||||
try:
|
||||
loader = TavilyLoader(
|
||||
urls=valid_urls,
|
||||
@@ -394,7 +388,7 @@ class SafeTavilyLoader(BaseLoader, RateLimitMixin, URLProcessingMixin):
|
||||
yield from loader.lazy_load()
|
||||
except Exception as e:
|
||||
if self.continue_on_failure:
|
||||
log.exception(f"Error extracting content from URLs: {e}")
|
||||
log.exception(f'Error extracting content from URLs: {e}')
|
||||
else:
|
||||
raise e
|
||||
|
||||
@@ -406,15 +400,15 @@ class SafeTavilyLoader(BaseLoader, RateLimitMixin, URLProcessingMixin):
|
||||
await self._safe_process_url(url)
|
||||
valid_urls.append(url)
|
||||
except Exception as e:
|
||||
log.warning(f"SSL verification failed for {url}: {str(e)}")
|
||||
log.warning(f'SSL verification failed for {url}: {str(e)}')
|
||||
if not self.continue_on_failure:
|
||||
raise e
|
||||
|
||||
if not valid_urls:
|
||||
if self.continue_on_failure:
|
||||
log.warning("No valid URLs to process after SSL verification")
|
||||
log.warning('No valid URLs to process after SSL verification')
|
||||
return
|
||||
raise ValueError("No valid URLs to process after SSL verification")
|
||||
raise ValueError('No valid URLs to process after SSL verification')
|
||||
|
||||
try:
|
||||
loader = TavilyLoader(
|
||||
@@ -427,7 +421,7 @@ class SafeTavilyLoader(BaseLoader, RateLimitMixin, URLProcessingMixin):
|
||||
yield document
|
||||
except Exception as e:
|
||||
if self.continue_on_failure:
|
||||
log.exception(f"Error loading URLs: {e}")
|
||||
log.exception(f'Error loading URLs: {e}')
|
||||
else:
|
||||
raise e
|
||||
|
||||
@@ -462,15 +456,15 @@ class SafePlaywrightURLLoader(PlaywrightURLLoader, RateLimitMixin, URLProcessing
|
||||
):
|
||||
"""Initialize with additional safety parameters and remote browser support."""
|
||||
|
||||
proxy_server = proxy.get("server") if proxy else None
|
||||
proxy_server = proxy.get('server') if proxy else None
|
||||
if trust_env and not proxy_server:
|
||||
env_proxies = urllib.request.getproxies()
|
||||
env_proxy_server = env_proxies.get("https") or env_proxies.get("http")
|
||||
env_proxy_server = env_proxies.get('https') or env_proxies.get('http')
|
||||
if env_proxy_server:
|
||||
if proxy:
|
||||
proxy["server"] = env_proxy_server
|
||||
proxy['server'] = env_proxy_server
|
||||
else:
|
||||
proxy = {"server": env_proxy_server}
|
||||
proxy = {'server': env_proxy_server}
|
||||
|
||||
# We'll set headless to False if using playwright_ws_url since it's handled by the remote browser
|
||||
super().__init__(
|
||||
@@ -504,14 +498,14 @@ class SafePlaywrightURLLoader(PlaywrightURLLoader, RateLimitMixin, URLProcessing
|
||||
page = browser.new_page()
|
||||
response = page.goto(url, timeout=self.playwright_timeout)
|
||||
if response is None:
|
||||
raise ValueError(f"page.goto() returned None for url {url}")
|
||||
raise ValueError(f'page.goto() returned None for url {url}')
|
||||
|
||||
text = self.evaluator.evaluate(page, browser, response)
|
||||
metadata = {"source": url}
|
||||
metadata = {'source': url}
|
||||
yield Document(page_content=text, metadata=metadata)
|
||||
except Exception as e:
|
||||
if self.continue_on_failure:
|
||||
log.exception(f"Error loading {url}: {e}")
|
||||
log.exception(f'Error loading {url}: {e}')
|
||||
continue
|
||||
raise e
|
||||
browser.close()
|
||||
@@ -525,9 +519,7 @@ class SafePlaywrightURLLoader(PlaywrightURLLoader, RateLimitMixin, URLProcessing
|
||||
if self.playwright_ws_url:
|
||||
browser = await p.chromium.connect(self.playwright_ws_url)
|
||||
else:
|
||||
browser = await p.chromium.launch(
|
||||
headless=self.headless, proxy=self.proxy
|
||||
)
|
||||
browser = await p.chromium.launch(headless=self.headless, proxy=self.proxy)
|
||||
|
||||
for url in self.urls:
|
||||
try:
|
||||
@@ -535,14 +527,14 @@ class SafePlaywrightURLLoader(PlaywrightURLLoader, RateLimitMixin, URLProcessing
|
||||
page = await browser.new_page()
|
||||
response = await page.goto(url, timeout=self.playwright_timeout)
|
||||
if response is None:
|
||||
raise ValueError(f"page.goto() returned None for url {url}")
|
||||
raise ValueError(f'page.goto() returned None for url {url}')
|
||||
|
||||
text = await self.evaluator.evaluate_async(page, browser, response)
|
||||
metadata = {"source": url}
|
||||
metadata = {'source': url}
|
||||
yield Document(page_content=text, metadata=metadata)
|
||||
except Exception as e:
|
||||
if self.continue_on_failure:
|
||||
log.exception(f"Error loading {url}: {e}")
|
||||
log.exception(f'Error loading {url}: {e}')
|
||||
continue
|
||||
raise e
|
||||
await browser.close()
|
||||
@@ -560,9 +552,7 @@ class SafeWebBaseLoader(WebBaseLoader):
|
||||
super().__init__(*args, **kwargs)
|
||||
self.trust_env = trust_env
|
||||
|
||||
async def _fetch(
|
||||
self, url: str, retries: int = 3, cooldown: int = 2, backoff: float = 1.5
|
||||
) -> str:
|
||||
async def _fetch(self, url: str, retries: int = 3, cooldown: int = 2, backoff: float = 1.5) -> str:
|
||||
async with aiohttp.ClientSession(trust_env=self.trust_env) as session:
|
||||
for i in range(retries):
|
||||
try:
|
||||
@@ -571,7 +561,7 @@ class SafeWebBaseLoader(WebBaseLoader):
|
||||
cookies=self.session.cookies.get_dict(),
|
||||
)
|
||||
if not self.session.verify:
|
||||
kwargs["ssl"] = False
|
||||
kwargs['ssl'] = False
|
||||
|
||||
async with session.get(
|
||||
url,
|
||||
@@ -585,16 +575,11 @@ class SafeWebBaseLoader(WebBaseLoader):
|
||||
if i == retries - 1:
|
||||
raise
|
||||
else:
|
||||
log.warning(
|
||||
f"Error fetching {url} with attempt "
|
||||
f"{i + 1}/{retries}: {e}. Retrying..."
|
||||
)
|
||||
log.warning(f'Error fetching {url} with attempt {i + 1}/{retries}: {e}. Retrying...')
|
||||
await asyncio.sleep(cooldown * backoff**i)
|
||||
raise ValueError("retry count exceeded")
|
||||
raise ValueError('retry count exceeded')
|
||||
|
||||
def _unpack_fetch_results(
|
||||
self, results: Any, urls: List[str], parser: Union[str, None] = None
|
||||
) -> List[Any]:
|
||||
def _unpack_fetch_results(self, results: Any, urls: List[str], parser: Union[str, None] = None) -> List[Any]:
|
||||
"""Unpack fetch results into BeautifulSoup objects."""
|
||||
from bs4 import BeautifulSoup
|
||||
|
||||
@@ -602,17 +587,15 @@ class SafeWebBaseLoader(WebBaseLoader):
|
||||
for i, result in enumerate(results):
|
||||
url = urls[i]
|
||||
if parser is None:
|
||||
if url.endswith(".xml"):
|
||||
parser = "xml"
|
||||
if url.endswith('.xml'):
|
||||
parser = 'xml'
|
||||
else:
|
||||
parser = self.default_parser
|
||||
self._check_parser(parser)
|
||||
final_results.append(BeautifulSoup(result, parser, **self.bs_kwargs))
|
||||
return final_results
|
||||
|
||||
async def ascrape_all(
|
||||
self, urls: List[str], parser: Union[str, None] = None
|
||||
) -> List[Any]:
|
||||
async def ascrape_all(self, urls: List[str], parser: Union[str, None] = None) -> List[Any]:
|
||||
"""Async fetch all urls, then return soups for all results."""
|
||||
results = await self.fetch_all(urls)
|
||||
return self._unpack_fetch_results(results, urls, parser=parser)
|
||||
@@ -630,22 +613,20 @@ class SafeWebBaseLoader(WebBaseLoader):
|
||||
yield Document(page_content=text, metadata=metadata)
|
||||
except Exception as e:
|
||||
# Log the error and continue with the next URL
|
||||
log.exception(f"Error loading {path}: {e}")
|
||||
log.exception(f'Error loading {path}: {e}')
|
||||
|
||||
async def alazy_load(self) -> AsyncIterator[Document]:
|
||||
"""Async lazy load text from the url(s) in web_path."""
|
||||
results = await self.ascrape_all(self.web_paths)
|
||||
for path, soup in zip(self.web_paths, results):
|
||||
text = soup.get_text(**self.bs_get_text_kwargs)
|
||||
metadata = {"source": path}
|
||||
if title := soup.find("title"):
|
||||
metadata["title"] = title.get_text()
|
||||
if description := soup.find("meta", attrs={"name": "description"}):
|
||||
metadata["description"] = description.get(
|
||||
"content", "No description found."
|
||||
)
|
||||
if html := soup.find("html"):
|
||||
metadata["language"] = html.get("lang", "No language found.")
|
||||
metadata = {'source': path}
|
||||
if title := soup.find('title'):
|
||||
metadata['title'] = title.get_text()
|
||||
if description := soup.find('meta', attrs={'name': 'description'}):
|
||||
metadata['description'] = description.get('content', 'No description found.')
|
||||
if html := soup.find('html'):
|
||||
metadata['language'] = html.get('lang', 'No language found.')
|
||||
yield Document(page_content=text, metadata=metadata)
|
||||
|
||||
async def aload(self) -> list[Document]:
|
||||
@@ -663,18 +644,18 @@ def get_web_loader(
|
||||
safe_urls = safe_validate_urls([urls] if isinstance(urls, str) else urls)
|
||||
|
||||
if not safe_urls:
|
||||
log.warning(f"All provided URLs were blocked or invalid: {urls}")
|
||||
log.warning(f'All provided URLs were blocked or invalid: {urls}')
|
||||
raise ValueError(ERROR_MESSAGES.INVALID_URL)
|
||||
|
||||
web_loader_args = {
|
||||
"web_paths": safe_urls,
|
||||
"verify_ssl": verify_ssl,
|
||||
"requests_per_second": requests_per_second,
|
||||
"continue_on_failure": True,
|
||||
"trust_env": trust_env,
|
||||
'web_paths': safe_urls,
|
||||
'verify_ssl': verify_ssl,
|
||||
'requests_per_second': requests_per_second,
|
||||
'continue_on_failure': True,
|
||||
'trust_env': trust_env,
|
||||
}
|
||||
|
||||
if WEB_LOADER_ENGINE.value == "" or WEB_LOADER_ENGINE.value == "safe_web":
|
||||
if WEB_LOADER_ENGINE.value == '' or WEB_LOADER_ENGINE.value == 'safe_web':
|
||||
WebLoaderClass = SafeWebBaseLoader
|
||||
|
||||
request_kwargs = {}
|
||||
@@ -685,42 +666,42 @@ def get_web_loader(
|
||||
timeout_value = None
|
||||
|
||||
if timeout_value:
|
||||
request_kwargs["timeout"] = timeout_value
|
||||
request_kwargs['timeout'] = timeout_value
|
||||
|
||||
if request_kwargs:
|
||||
web_loader_args["requests_kwargs"] = request_kwargs
|
||||
web_loader_args['requests_kwargs'] = request_kwargs
|
||||
|
||||
if WEB_LOADER_ENGINE.value == "playwright":
|
||||
if WEB_LOADER_ENGINE.value == 'playwright':
|
||||
WebLoaderClass = SafePlaywrightURLLoader
|
||||
web_loader_args["playwright_timeout"] = PLAYWRIGHT_TIMEOUT.value
|
||||
web_loader_args['playwright_timeout'] = PLAYWRIGHT_TIMEOUT.value
|
||||
if PLAYWRIGHT_WS_URL.value:
|
||||
web_loader_args["playwright_ws_url"] = PLAYWRIGHT_WS_URL.value
|
||||
web_loader_args['playwright_ws_url'] = PLAYWRIGHT_WS_URL.value
|
||||
|
||||
if WEB_LOADER_ENGINE.value == "firecrawl":
|
||||
if WEB_LOADER_ENGINE.value == 'firecrawl':
|
||||
WebLoaderClass = SafeFireCrawlLoader
|
||||
web_loader_args["api_key"] = FIRECRAWL_API_KEY.value
|
||||
web_loader_args["api_url"] = FIRECRAWL_API_BASE_URL.value
|
||||
web_loader_args['api_key'] = FIRECRAWL_API_KEY.value
|
||||
web_loader_args['api_url'] = FIRECRAWL_API_BASE_URL.value
|
||||
if FIRECRAWL_TIMEOUT.value:
|
||||
try:
|
||||
web_loader_args["timeout"] = int(FIRECRAWL_TIMEOUT.value)
|
||||
web_loader_args['timeout'] = int(FIRECRAWL_TIMEOUT.value)
|
||||
except ValueError:
|
||||
pass
|
||||
|
||||
if WEB_LOADER_ENGINE.value == "tavily":
|
||||
if WEB_LOADER_ENGINE.value == 'tavily':
|
||||
WebLoaderClass = SafeTavilyLoader
|
||||
web_loader_args["api_key"] = TAVILY_API_KEY.value
|
||||
web_loader_args["extract_depth"] = TAVILY_EXTRACT_DEPTH.value
|
||||
web_loader_args['api_key'] = TAVILY_API_KEY.value
|
||||
web_loader_args['extract_depth'] = TAVILY_EXTRACT_DEPTH.value
|
||||
|
||||
if WEB_LOADER_ENGINE.value == "external":
|
||||
if WEB_LOADER_ENGINE.value == 'external':
|
||||
WebLoaderClass = ExternalWebLoader
|
||||
web_loader_args["external_url"] = EXTERNAL_WEB_LOADER_URL.value
|
||||
web_loader_args["external_api_key"] = EXTERNAL_WEB_LOADER_API_KEY.value
|
||||
web_loader_args['external_url'] = EXTERNAL_WEB_LOADER_URL.value
|
||||
web_loader_args['external_api_key'] = EXTERNAL_WEB_LOADER_API_KEY.value
|
||||
|
||||
if WebLoaderClass:
|
||||
web_loader = WebLoaderClass(**web_loader_args)
|
||||
|
||||
log.debug(
|
||||
"Using WEB_LOADER_ENGINE %s for %s URLs",
|
||||
'Using WEB_LOADER_ENGINE %s for %s URLs',
|
||||
web_loader.__class__.__name__,
|
||||
len(safe_urls),
|
||||
)
|
||||
@@ -728,6 +709,6 @@ def get_web_loader(
|
||||
return web_loader
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Invalid WEB_LOADER_ENGINE: {WEB_LOADER_ENGINE.value}. "
|
||||
f'Invalid WEB_LOADER_ENGINE: {WEB_LOADER_ENGINE.value}. '
|
||||
"Please set it to 'safe_web', 'playwright', 'firecrawl', or 'tavily'."
|
||||
)
|
||||
|
||||
@@ -41,29 +41,29 @@ def search_yacy(
|
||||
yacy_auth = HTTPDigestAuth(username, password)
|
||||
|
||||
params = {
|
||||
"query": query,
|
||||
"contentdom": "text",
|
||||
"resource": "global",
|
||||
"maximumRecords": count,
|
||||
"nav": "none",
|
||||
'query': query,
|
||||
'contentdom': 'text',
|
||||
'resource': 'global',
|
||||
'maximumRecords': count,
|
||||
'nav': 'none',
|
||||
}
|
||||
|
||||
# Check if provided a json API URL
|
||||
if not query_url.endswith("yacysearch.json"):
|
||||
if not query_url.endswith('yacysearch.json'):
|
||||
# Strip all query parameters from the URL
|
||||
query_url = query_url.rstrip("/") + "/yacysearch.json"
|
||||
query_url = query_url.rstrip('/') + '/yacysearch.json'
|
||||
|
||||
log.debug(f"searching {query_url}")
|
||||
log.debug(f'searching {query_url}')
|
||||
|
||||
response = requests.get(
|
||||
query_url,
|
||||
auth=yacy_auth,
|
||||
headers={
|
||||
"User-Agent": "Open WebUI (https://github.com/open-webui/open-webui) RAG Bot",
|
||||
"Accept": "text/html",
|
||||
"Accept-Encoding": "gzip, deflate",
|
||||
"Accept-Language": "en-US,en;q=0.5",
|
||||
"Connection": "keep-alive",
|
||||
'User-Agent': 'Open WebUI (https://github.com/open-webui/open-webui) RAG Bot',
|
||||
'Accept': 'text/html',
|
||||
'Accept-Encoding': 'gzip, deflate',
|
||||
'Accept-Language': 'en-US,en;q=0.5',
|
||||
'Connection': 'keep-alive',
|
||||
},
|
||||
params=params,
|
||||
)
|
||||
@@ -71,15 +71,15 @@ def search_yacy(
|
||||
response.raise_for_status() # Raise an exception for HTTP errors.
|
||||
|
||||
json_response = response.json()
|
||||
results = json_response.get("channels", [{}])[0].get("items", [])
|
||||
sorted_results = sorted(results, key=lambda x: x.get("ranking", 0), reverse=True)
|
||||
results = json_response.get('channels', [{}])[0].get('items', [])
|
||||
sorted_results = sorted(results, key=lambda x: x.get('ranking', 0), reverse=True)
|
||||
if filter_list:
|
||||
sorted_results = get_filtered_results(sorted_results, filter_list)
|
||||
return [
|
||||
SearchResult(
|
||||
link=result["link"],
|
||||
title=result.get("title"),
|
||||
snippet=result.get("description"),
|
||||
link=result['link'],
|
||||
title=result.get('title'),
|
||||
snippet=result.get('description'),
|
||||
)
|
||||
for result in sorted_results[:count]
|
||||
]
|
||||
|
||||
@@ -20,14 +20,14 @@ log = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def xml_element_contents_to_string(element: Element) -> str:
|
||||
buffer = [element.text if element.text else ""]
|
||||
buffer = [element.text if element.text else '']
|
||||
|
||||
for child in element:
|
||||
buffer.append(xml_element_contents_to_string(child))
|
||||
|
||||
buffer.append(element.tail if element.tail else "")
|
||||
buffer.append(element.tail if element.tail else '')
|
||||
|
||||
return "".join(buffer)
|
||||
return ''.join(buffer)
|
||||
|
||||
|
||||
def search_yandex(
|
||||
@@ -42,42 +42,38 @@ def search_yandex(
|
||||
) -> List[SearchResult]:
|
||||
try:
|
||||
headers = {
|
||||
"User-Agent": "Open WebUI (https://github.com/open-webui/open-webui) RAG Bot",
|
||||
"Authorization": f"Api-Key {yandex_search_api_key}",
|
||||
'User-Agent': 'Open WebUI (https://github.com/open-webui/open-webui) RAG Bot',
|
||||
'Authorization': f'Api-Key {yandex_search_api_key}',
|
||||
}
|
||||
|
||||
if user is not None:
|
||||
headers = include_user_info_headers(headers, user)
|
||||
|
||||
chat_id = getattr(request.state, "chat_id", None)
|
||||
chat_id = getattr(request.state, 'chat_id', None)
|
||||
if chat_id:
|
||||
headers[FORWARD_SESSION_INFO_HEADER_CHAT_ID] = str(chat_id)
|
||||
|
||||
payload = {} if yandex_search_config == "" else json.loads(yandex_search_config)
|
||||
payload = {} if yandex_search_config == '' else json.loads(yandex_search_config)
|
||||
|
||||
if type(payload.get("query", None)) != dict:
|
||||
payload["query"] = {}
|
||||
if type(payload.get('query', None)) != dict:
|
||||
payload['query'] = {}
|
||||
|
||||
if "searchType" not in payload["query"]:
|
||||
payload["query"]["searchType"] = "SEARCH_TYPE_RU"
|
||||
if 'searchType' not in payload['query']:
|
||||
payload['query']['searchType'] = 'SEARCH_TYPE_RU'
|
||||
|
||||
payload["query"]["queryText"] = query
|
||||
payload['query']['queryText'] = query
|
||||
|
||||
if type(payload.get("groupSpec", None)) != dict:
|
||||
payload["groupSpec"] = {}
|
||||
if type(payload.get('groupSpec', None)) != dict:
|
||||
payload['groupSpec'] = {}
|
||||
|
||||
if "groupMode" not in payload["groupSpec"]:
|
||||
payload["groupSpec"]["groupMode"] = "GROUP_MODE_DEEP"
|
||||
if 'groupMode' not in payload['groupSpec']:
|
||||
payload['groupSpec']['groupMode'] = 'GROUP_MODE_DEEP'
|
||||
|
||||
payload["groupSpec"]["groupsOnPage"] = count
|
||||
payload["groupSpec"]["docsInGroup"] = 1
|
||||
payload['groupSpec']['groupsOnPage'] = count
|
||||
payload['groupSpec']['docsInGroup'] = 1
|
||||
|
||||
response = requests.post(
|
||||
(
|
||||
"https://searchapi.api.cloud.yandex.net/v2/web/search"
|
||||
if yandex_search_url == ""
|
||||
else yandex_search_url
|
||||
),
|
||||
('https://searchapi.api.cloud.yandex.net/v2/web/search' if yandex_search_url == '' else yandex_search_url),
|
||||
headers=headers,
|
||||
json=payload,
|
||||
)
|
||||
@@ -85,29 +81,21 @@ def search_yandex(
|
||||
response.raise_for_status()
|
||||
|
||||
response_body = response.json()
|
||||
if "rawData" not in response_body:
|
||||
raise Exception(f"No `rawData` in response body: {response_body}")
|
||||
if 'rawData' not in response_body:
|
||||
raise Exception(f'No `rawData` in response body: {response_body}')
|
||||
|
||||
search_result_body_bytes = base64.decodebytes(
|
||||
bytes(response_body["rawData"], "utf-8")
|
||||
)
|
||||
search_result_body_bytes = base64.decodebytes(bytes(response_body['rawData'], 'utf-8'))
|
||||
|
||||
doc_root = ET.parse(io.BytesIO(search_result_body_bytes))
|
||||
|
||||
results = []
|
||||
|
||||
for group in doc_root.findall("response/results/grouping/group"):
|
||||
for group in doc_root.findall('response/results/grouping/group'):
|
||||
results.append(
|
||||
{
|
||||
"url": xml_element_contents_to_string(group.find("doc/url")).strip(
|
||||
"\n"
|
||||
),
|
||||
"title": xml_element_contents_to_string(
|
||||
group.find("doc/title")
|
||||
).strip("\n"),
|
||||
"snippet": xml_element_contents_to_string(
|
||||
group.find("doc/passages/passage")
|
||||
),
|
||||
'url': xml_element_contents_to_string(group.find('doc/url')).strip('\n'),
|
||||
'title': xml_element_contents_to_string(group.find('doc/title')).strip('\n'),
|
||||
'snippet': xml_element_contents_to_string(group.find('doc/passages/passage')),
|
||||
}
|
||||
)
|
||||
|
||||
@@ -115,49 +103,47 @@ def search_yandex(
|
||||
|
||||
results = [
|
||||
SearchResult(
|
||||
link=result.get("url"),
|
||||
title=result.get("title"),
|
||||
snippet=result.get("snippet"),
|
||||
link=result.get('url'),
|
||||
title=result.get('title'),
|
||||
snippet=result.get('snippet'),
|
||||
)
|
||||
for result in results[:count]
|
||||
]
|
||||
|
||||
log.info(f"Yandex search results: {results}")
|
||||
log.info(f'Yandex search results: {results}')
|
||||
|
||||
return results
|
||||
except Exception as e:
|
||||
log.error(f"Error in search: {e}")
|
||||
log.error(f'Error in search: {e}')
|
||||
|
||||
return []
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
if __name__ == '__main__':
|
||||
from starlette.datastructures import Headers
|
||||
from fastapi import FastAPI
|
||||
|
||||
result = search_yandex(
|
||||
Request(
|
||||
{
|
||||
"type": "http",
|
||||
"asgi.version": "3.0",
|
||||
"asgi.spec_version": "2.0",
|
||||
"method": "GET",
|
||||
"path": "/internal",
|
||||
"query_string": b"",
|
||||
"headers": Headers({}).raw,
|
||||
"client": ("127.0.0.1", 12345),
|
||||
"server": ("127.0.0.1", 80),
|
||||
"scheme": "http",
|
||||
"app": FastAPI(),
|
||||
'type': 'http',
|
||||
'asgi.version': '3.0',
|
||||
'asgi.spec_version': '2.0',
|
||||
'method': 'GET',
|
||||
'path': '/internal',
|
||||
'query_string': b'',
|
||||
'headers': Headers({}).raw,
|
||||
'client': ('127.0.0.1', 12345),
|
||||
'server': ('127.0.0.1', 80),
|
||||
'scheme': 'http',
|
||||
'app': FastAPI(),
|
||||
},
|
||||
None,
|
||||
),
|
||||
os.environ.get("YANDEX_WEB_SEARCH_URL", ""),
|
||||
os.environ.get("YANDEX_WEB_SEARCH_API_KEY", ""),
|
||||
os.environ.get(
|
||||
"YANDEX_WEB_SEARCH_CONFIG", '{"query": {"searchType": "SEARCH_TYPE_COM"}}'
|
||||
),
|
||||
"TOP movies of the past year",
|
||||
os.environ.get('YANDEX_WEB_SEARCH_URL', ''),
|
||||
os.environ.get('YANDEX_WEB_SEARCH_API_KEY', ''),
|
||||
os.environ.get('YANDEX_WEB_SEARCH_CONFIG', '{"query": {"searchType": "SEARCH_TYPE_COM"}}'),
|
||||
'TOP movies of the past year',
|
||||
3,
|
||||
)
|
||||
|
||||
|
||||
@@ -12,7 +12,7 @@ def search_youcom(
|
||||
query: str,
|
||||
count: int,
|
||||
filter_list: Optional[List[str]] = None,
|
||||
language: str = "EN",
|
||||
language: str = 'EN',
|
||||
) -> List[SearchResult]:
|
||||
"""Search using You.com's YDC Index API and return the results as a list of SearchResult objects.
|
||||
|
||||
@@ -23,30 +23,30 @@ def search_youcom(
|
||||
filter_list (list[str], optional): Domain filter list
|
||||
language (str): Language code for search results (default: "EN")
|
||||
"""
|
||||
url = "https://ydc-index.io/v1/search"
|
||||
url = 'https://ydc-index.io/v1/search'
|
||||
headers = {
|
||||
"Accept": "application/json",
|
||||
"X-API-KEY": api_key,
|
||||
'Accept': 'application/json',
|
||||
'X-API-KEY': api_key,
|
||||
}
|
||||
params = {
|
||||
"query": query,
|
||||
"count": count,
|
||||
"language": language,
|
||||
'query': query,
|
||||
'count': count,
|
||||
'language': language,
|
||||
}
|
||||
|
||||
response = requests.get(url, headers=headers, params=params)
|
||||
response.raise_for_status()
|
||||
|
||||
json_response = response.json()
|
||||
results = json_response.get("results", {}).get("web", [])
|
||||
results = json_response.get('results', {}).get('web', [])
|
||||
|
||||
if filter_list:
|
||||
results = get_filtered_results(results, filter_list)
|
||||
|
||||
return [
|
||||
SearchResult(
|
||||
link=result["url"],
|
||||
title=result.get("title"),
|
||||
link=result['url'],
|
||||
title=result.get('title'),
|
||||
snippet=_build_snippet(result),
|
||||
)
|
||||
for result in results[:count]
|
||||
@@ -62,12 +62,12 @@ def _build_snippet(result: dict) -> str:
|
||||
"""
|
||||
parts: list[str] = []
|
||||
|
||||
description = result.get("description")
|
||||
description = result.get('description')
|
||||
if description:
|
||||
parts.append(description)
|
||||
|
||||
snippets = result.get("snippets")
|
||||
snippets = result.get('snippets')
|
||||
if snippets and isinstance(snippets, list):
|
||||
parts.extend(snippets)
|
||||
|
||||
return "\n\n".join(parts)
|
||||
return '\n\n'.join(parts)
|
||||
|
||||
Reference in New Issue
Block a user