argo-search 1.0.1

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -0,0 +1,386 @@
1
+ #!/usr/bin/env python3
2
+ """engines.py — Unified Search v2 引擎适配层(精简版,< 400 行)
3
+
4
+ 配置驱动 + 声明式 output_map 字段提取 + 通用 parser 兜底。
5
+ 支持 cli / http(GET/POST) 类型,所有异常吞没返回 []。
6
+ """
7
+
8
+ from __future__ import annotations
9
+
10
+ import functools
11
+ import json
12
+ import logging
13
+ import os
14
+ import re
15
+ import subprocess
16
+ import sys
17
+ import time
18
+ import urllib.error
19
+ import urllib.request
20
+ from pathlib import Path
21
+ from typing import Any, Callable
22
+
23
+ try:
24
+ from config import load_config, get_engines
25
+ except ImportError:
26
+ sys.path.insert(0, str(Path(__file__).parent))
27
+ from config import load_config, get_engines
28
+
29
+ logger = logging.getLogger("unified_search.engines")
30
+ if not logger.handlers:
31
+ logger.setLevel(logging.WARNING)
32
+ logger.addHandler(logging.StreamHandler(sys.stderr))
33
+
34
+
35
+ def safe_search(fn: Callable) -> Callable:
36
+ """统一错误处理装饰器 — 所有异常返回 []。"""
37
+ @functools.wraps(fn)
38
+ def wrapper(*args, **kwargs) -> list[dict[str, Any]]:
39
+ name = fn.__name__.replace("_engine", "").strip("_")
40
+ try:
41
+ return fn(*args, **kwargs)
42
+ except subprocess.TimeoutExpired:
43
+ logger.warning(f"引擎 {name} 超时")
44
+ except FileNotFoundError as e:
45
+ logger.warning(f"引擎 {name} 命令不存在: {e}")
46
+ except Exception as e:
47
+ logger.error(f"引擎 {name} 异常: {type(e).__name__}: {e}", exc_info=True)
48
+ return []
49
+ return wrapper
50
+
51
+
52
+ def _run(cmd: list[str], timeout: float = 8, engine_name: str = "?") -> str:
53
+ """执行命令,超时/异常不抛。"""
54
+ try:
55
+ r = subprocess.run(cmd, capture_output=True, text=True, timeout=timeout)
56
+ if r.returncode == 0:
57
+ return r.stdout
58
+ tail = (r.stderr or "").strip()[:200]
59
+ logger.warning(f"引擎 {engine_name} 失败 (rc={r.returncode}): {tail}")
60
+ return r.stdout if r.stdout.strip() else ""
61
+ except subprocess.TimeoutExpired:
62
+ logger.warning(f"引擎 {engine_name} 超时 (>{timeout}s)")
63
+ except FileNotFoundError as e:
64
+ logger.error(f"引擎 {engine_name} CLI 缺失: {e}")
65
+ except Exception as e:
66
+ logger.error(f"引擎 {engine_name} 异常: {type(e).__name__}: {e}")
67
+ return ""
68
+
69
+
70
+ def _resolve(template: list[str] | str, query: str, n: int, **extra: Any) -> list[str] | str:
71
+ """替换模板占位符。"""
72
+ if isinstance(template, list):
73
+ return [_resolve(item, query, n, **extra) for item in template]
74
+ s = template.replace("{query}", query).replace("{n}", str(n))
75
+ s = s.replace("{TIMESTAMP}", str(int(time.time())))
76
+ for key, val in extra.items():
77
+ s = s.replace(f"{{{key}}}", str(val))
78
+ if s.startswith("~"):
79
+ s = str(Path.home() / s[1:])
80
+ return re.sub(r"\{([A-Z_][A-Z0-9_]*)\}", lambda m: os.environ.get(m.group(1), m.group(0)), s)
81
+
82
+
83
+ def _extract_items(data: Any, path: str) -> list[dict]:
84
+ """从嵌套 dict 按路径提取列表。"""
85
+ obj = data
86
+ for part in path.split("."):
87
+ if isinstance(obj, dict):
88
+ obj = obj.get(part, [])
89
+ else:
90
+ return []
91
+ return obj if isinstance(obj, list) else []
92
+
93
+
94
+ def _make_field_parser(path: str, fields: dict[str, str]) -> Callable:
95
+ """构造声明式 parser。"""
96
+ def parser(data: Any) -> list[dict[str, Any]]:
97
+ items = _extract_items(data, path) if isinstance(data, dict) else []
98
+ results = []
99
+ for item in items:
100
+ if not isinstance(item, dict):
101
+ continue
102
+ r = {ok: iv[:300] if ok == "snippet" and isinstance(iv, str) else iv
103
+ for ok, ik in fields.items() if (iv := item.get(ik, ""))}
104
+ if r.get("title") or r.get("url"):
105
+ results.append(r)
106
+ return results[:10]
107
+ return parser
108
+
109
+
110
+ def _build_cli_engine(spec: dict[str, Any]) -> Any:
111
+ cmd_template = spec.get("cmd", [])
112
+ search_args = spec.get("search_args", [])
113
+ env_overrides = spec.get("env", {})
114
+
115
+ @safe_search
116
+ def _engine(query: str, n: int = 5, timeout: float = 8, mode: str = "fast", **kwargs) -> list[dict[str, Any]]:
117
+ cmd = _resolve(cmd_template, query, n, mode=mode)
118
+ args = _resolve(search_args, query, n, mode=mode)
119
+ if not cmd:
120
+ return []
121
+ env = os.environ.copy()
122
+ env.update(env_overrides)
123
+ return _parse_text_output(_run(cmd + args, timeout=timeout, engine_name=spec.get("_name", "cli")),
124
+ spec.get("_name", "cli"))
125
+ return _engine
126
+
127
+
128
+ def _build_http_engine(spec: dict[str, Any]) -> Any:
129
+ """统一 HTTP 引擎构造(GET/POST)。"""
130
+ url_template = spec.get("url", "")
131
+ headers = spec.get("headers", {"Content-Type": "application/json"})
132
+ query_param = spec.get("query_param", "q")
133
+ fmt = spec.get("format", "")
134
+ timeout = spec.get("timeout", 8)
135
+ extra_params = spec.get("extra_params", {})
136
+ output_map = spec.get("output_map", {})
137
+ is_get = spec.get("method", "GET") == "GET"
138
+ body_template = spec.get("body", {})
139
+
140
+ @safe_search
141
+ def _engine(query: str, n: int = 5, _timeout: float | None = None, depth: str = "fast", **kwargs) -> list[dict[str, Any]]:
142
+ to = _timeout or timeout
143
+ import urllib.parse as up
144
+
145
+ if is_get:
146
+ resolved_url = _resolve(url_template, query, n)
147
+ separator = "&" if "?" in resolved_url else "?"
148
+ full_url = f"{resolved_url}{separator}{query_param}={up.quote(query)}"
149
+ if fmt:
150
+ full_url += f"&format={fmt}"
151
+ for k, v in extra_params.items():
152
+ full_url += f"&{k}={up.quote(_resolve(str(v), query, n))}"
153
+ req = urllib.request.Request(full_url, headers={k: _resolve(v, query, n) for k, v in headers.items()})
154
+ else:
155
+ body: dict[str, Any] = {}
156
+ for k, v in body_template.items():
157
+ resolved = _resolve(str(v), query, n)
158
+ if k == "search_depth":
159
+ body[k] = depth
160
+ elif resolved.lower() == "true":
161
+ body[k] = True
162
+ elif resolved.lower() == "false":
163
+ body[k] = False
164
+ else:
165
+ try:
166
+ body[k] = int(resolved)
167
+ except ValueError:
168
+ try:
169
+ body[k] = float(resolved)
170
+ except ValueError:
171
+ body[k] = resolved
172
+ req = urllib.request.Request(url_template, data=json.dumps(body).encode("utf-8"),
173
+ headers={k: _resolve(v, query, n) for k, v in headers.items()})
174
+
175
+ try:
176
+ with urllib.request.urlopen(req, timeout=to) as resp:
177
+ raw = resp.read().decode("utf-8")
178
+ if fmt == "xml":
179
+ return _parse_xml(raw, spec.get("_name", ""))
180
+ data = json.loads(raw)
181
+ if output_map:
182
+ return _make_field_parser(output_map.get("items", ""), {
183
+ "title": output_map.get("item_title", "title"),
184
+ "url": output_map.get("item_url", "url"),
185
+ "snippet": output_map.get("item_summary", "snippet"),
186
+ "source": output_map.get("item_source", "source"),
187
+ })(data)
188
+ return _parse_generic(data, spec.get("_name", ""))
189
+ except (urllib.error.URLError, urllib.error.HTTPError, Exception) as e:
190
+ logger.warning(f"HTTP 引擎失败: {e}")
191
+ return []
192
+ return _engine
193
+
194
+
195
+ _BUILDERS = {"cli": _build_cli_engine, "http": _build_http_engine}
196
+
197
+
198
+ # ── 通用解析器 ─────────────────────────────────────────────────────────────────
199
+
200
+ def _parse_text_output(text: str, engine_name: str) -> list[dict[str, Any]]:
201
+ """通用 CLI 文本解析:优先 JSON,其次结构化文本。"""
202
+ try:
203
+ data = json.loads(text.strip())
204
+ if isinstance(data, list):
205
+ return [{"title": i.get("title", ""), "url": i.get("url", ""),
206
+ "snippet": i.get("snippet", i.get("content", ""))[:300],
207
+ "source": engine_name} for i in data if isinstance(i, dict)]
208
+ if isinstance(data, dict):
209
+ items = data.get("results", data.get("items", data.get("data", [])))
210
+ if isinstance(items, list):
211
+ return [{"title": i.get("title", ""), "url": i.get("url", ""),
212
+ "snippet": i.get("snippet", i.get("content", ""))[:300],
213
+ "source": engine_name} for i in items if isinstance(i, dict)]
214
+ except (json.JSONDecodeError, ValueError):
215
+ pass
216
+
217
+ results, cur = [], {}
218
+ seen_url = False
219
+ for line in text.split("\n"):
220
+ s = line.strip()
221
+ if s.startswith("### "):
222
+ if cur:
223
+ results.append(cur)
224
+ cur = {"title": re.sub(r'^\d+\.\s*', '', s[4:].strip()), "source": engine_name,
225
+ "score": max(1.0 - len(results) * 0.1, 0.1)}
226
+ seen_url = False
227
+ elif s.startswith("- **URL**: ") and cur:
228
+ cur["url"] = s[11:].strip()
229
+ seen_url = True
230
+ elif s.startswith("- ") and not s.startswith("- **") and seen_url and cur:
231
+ cur["snippet"] = " ".join(s[2:].strip().split())[:300]
232
+ seen_url = False
233
+ if cur:
234
+ results.append(cur)
235
+ return results[:10]
236
+
237
+
238
+ def _parse_xml(text: str, engine_name: str) -> list[dict[str, Any]]:
239
+ """解析 Atom XML(arXiv 等)。"""
240
+ import xml.etree.ElementTree as ET
241
+ results = []
242
+ try:
243
+ root = ET.fromstring(text)
244
+ ns = {"atom": "http://www.w3.org/2005/Atom"}
245
+ entries = root.findall(".//atom:entry", ns) or root.findall(".//{http://www.w3.org/2005/Atom}entry")
246
+ for entry in entries:
247
+ title = entry.findtext("atom:title", "", ns).strip().replace("\n", " ")[:200]
248
+ summary = entry.findtext("atom:summary", "", ns).strip().replace("\n", " ")[:300]
249
+ entry_id = entry.findtext("atom:id", "", ns)
250
+ url = entry_id
251
+ for link in entry.findall("atom:link", ns):
252
+ if link.get("title") == "pdf":
253
+ url = link.get("href", url)
254
+ break
255
+ if title:
256
+ results.append({"title": title, "url": url, "snippet": summary, "source": engine_name})
257
+ except ET.ParseError:
258
+ pass
259
+ return results
260
+
261
+
262
+ def _parse_generic(data: dict[str, Any], engine_name: str = "?") -> list[dict[str, Any]]:
263
+ """通用 JSON 解析:自动探测常见字段。"""
264
+ items = None
265
+ for key in ["results", "items", "data", "works", "message.items"]:
266
+ if "." in key:
267
+ parts = key.split(".")
268
+ obj = data
269
+ for p in parts:
270
+ obj = obj.get(p, {}) if isinstance(obj, dict) else {}
271
+ if isinstance(obj, list):
272
+ items = obj
273
+ break
274
+ elif isinstance(data, dict) and key in data and isinstance(data[key], list):
275
+ items = data[key]
276
+ break
277
+
278
+ if items is None and isinstance(data, dict):
279
+ for v in data.values():
280
+ if isinstance(v, dict):
281
+ for key in ["results", "items", "value"]:
282
+ if key in v and isinstance(v[key], list):
283
+ items = v[key]
284
+ break
285
+ if items:
286
+ break
287
+
288
+ if not items or not isinstance(items, list):
289
+ return []
290
+
291
+ results = []
292
+ for i in items:
293
+ if not isinstance(i, dict):
294
+ continue
295
+ title = i.get("title", "")
296
+ if isinstance(title, list):
297
+ title = title[0] if title else ""
298
+ url = i.get("url", i.get("URL", i.get("html_url", "")))
299
+ snippet = (i.get("snippet", i.get("content", i.get("summary", i.get("description", "")))))[:300]
300
+ score = i.get("score", i.get("relevance_score", 0.5))
301
+ results.append({"title": str(title)[:200], "url": str(url),
302
+ "snippet": str(snippet), "score": score, "source": engine_name})
303
+ return results[:10]
304
+
305
+
306
+ # ── 引擎注册表 ─────────────────────────────────────────────────────────────────
307
+
308
+ _engine_registry: dict[str, Any] = {}
309
+ _engine_registry_loaded = False
310
+
311
+
312
+ def _load_registry():
313
+ global _engine_registry, _engine_registry_loaded
314
+ if _engine_registry_loaded:
315
+ return
316
+ cfg = load_config()
317
+ engines = get_engines(cfg)
318
+ registry = {}
319
+ for name, spec in engines.items():
320
+ spec = dict(spec)
321
+ spec["_name"] = name
322
+ builder = _BUILDERS.get(spec.get("type", "cli"))
323
+ if builder:
324
+ registry[name] = builder(spec)
325
+ else:
326
+ logger.warning(f"未知引擎类型: {spec.get('type')} (引擎 {name})")
327
+ _engine_registry = registry
328
+ _engine_registry_loaded = True
329
+
330
+
331
+ def get_registry() -> dict[str, Any]:
332
+ _load_registry()
333
+ return _engine_registry
334
+
335
+
336
+ def available_engines() -> list[str]:
337
+ return sorted(get_registry().keys())
338
+
339
+
340
+ def search(query: str, engine: str, n: int = 5, timeout: float = 8, depth: str = "fast", mode: str = "fast") -> list[dict[str, Any]]:
341
+ """统一引擎调用入口;失败返回空 list,不抛异常。"""
342
+ registry = get_registry()
343
+ fn = registry.get(engine)
344
+ if not fn:
345
+ logger.warning(f"未知引擎: {engine}")
346
+ return []
347
+ t0 = time.time()
348
+ try:
349
+ results = fn(query, n, timeout, depth=depth, mode=mode)
350
+ except TypeError:
351
+ try:
352
+ results = fn(query, n, timeout)
353
+ except Exception as e:
354
+ logger.error(f"引擎 {engine} 异常: {type(e).__name__}: {e}")
355
+ results = []
356
+ except Exception as e:
357
+ logger.error(f"引擎 {engine} 异常: {type(e).__name__}: {e}")
358
+ results = []
359
+ elapsed = time.time() - t0
360
+ if results and isinstance(results, list):
361
+ for r in results:
362
+ if isinstance(r, dict) and "error" not in r:
363
+ r["_engine"] = engine
364
+ r["_elapsed"] = round(elapsed, 3)
365
+ return results if isinstance(results, list) else []
366
+
367
+
368
+ def _cli():
369
+ import argparse
370
+ parser = argparse.ArgumentParser(description="引擎适配层调试")
371
+ parser.add_argument("query", nargs="?")
372
+ parser.add_argument("--engine", "-e", default="anysearch")
373
+ parser.add_argument("-n", type=int, default=5)
374
+ parser.add_argument("--timeout", "-t", type=float, default=8)
375
+ parser.add_argument("--list", action="store_true")
376
+ args = parser.parse_args()
377
+ if args.list:
378
+ print(json.dumps(available_engines(), ensure_ascii=False, indent=2))
379
+ return
380
+ if not args.query:
381
+ parser.error("必须提供 query")
382
+ print(json.dumps(search(args.query, args.engine, args.n, args.timeout), ensure_ascii=False, indent=2))
383
+
384
+
385
+ if __name__ == "__main__":
386
+ _cli()