Coverage for open_webui/retrieval/web/staan.py: 21%

20 statements  

« prev     ^ index     » next       coverage.py v7.15.2, created at 2026-10-07 05:07 +0000

1from __future__ import annotations 

2 

3import requests 

4from open_webui.retrieval.web.main import SearchResult, get_filtered_results 

5 

6 

7def search_staan( 

8 api_key: str, 

9 query: str, 

10 count: int, 

11 filter_list: list[str] | None = None, 

12 market: str | None = None, 

13 max_snippets: int | None = None, 

14) -> list[SearchResult]: 

15 """Search using Staan's Web Search API and return the results as a list of SearchResult objects. 

16 

17 Args: 

18 api_key (str): A Staan API key 

19 query (str): The query to search for 

20 count (int): The maximum number of results to return 

21 filter_list (list[str] | None): The domains to allow or block 

22 market (str | None): The market to search in, e.g. 'en-us' 

23 max_snippets (int | None): The maximum extra snippets to request per result 

24 

25 Returns: 

26 A list of SearchResult objects. 

27 """ 

28 url = 'https://api.staan.ai/v2/search/web' 

29 headers = { 

30 'Accept': 'application/json', 

31 'Authorization': f'Bearer {api_key}', 

32 } 

33 params = {'q': query, 'market': market} 

34 

35 if max_snippets: 

36 params['extra_snippets'] = 'true' 

37 params['max_snippets'] = max_snippets 

38 

39 response = requests.get(url, headers=headers, params=params) 

40 response.raise_for_status() 

41 

42 results = response.json().get('web', {}).get('results', []) 

43 if filter_list: 

44 results = get_filtered_results(results, filter_list) 

45 

46 return [ 

47 SearchResult( 

48 link=result.get('url', ''), 

49 title=result.get('title'), 

50 snippet=_build_snippet(result), 

51 ) 

52 for result in results[:count] 

53 ] 

54 

55 

56def _build_snippet(result: dict) -> str: 

57 """Combine the snippet and the extra snippets list into a single string.""" 

58 parts = [result.get('snippet')] 

59 parts.extend(extra.get('chunk') for extra in result.get('extra_snippets', [])) 

60 return '\n\n'.join(part for part in parts if part)