Edit on GitHub

agent_search_gateway.cli

Thin command-line client for the local foreground gateway daemon.

  1"""Thin command-line client for the local foreground gateway daemon."""
  2
  3import argparse
  4import asyncio
  5import os
  6import signal
  7import sys
  8from collections.abc import Awaitable, Callable, Mapping
  9from pathlib import Path
 10from typing import Protocol, TextIO
 11
 12from .daemon import ForegroundDaemon
 13from .doctor import DoctorReport, render_doctor, run_doctor
 14from .errors import DaemonUnavailable, ErrorCode, GatewayError, InputFailure
 15from .models import (
 16    ErrorResponse,
 17    KeywordSearchRequest,
 18    LLMSearchRequest,
 19    PaperSearchRequest,
 20    Request,
 21    Response,
 22    ShutdownRequest,
 23    SuccessResponse,
 24    URLFetchRequest,
 25)
 26from .observability import DebugLoggingSession, configure_debug_logging
 27from .paths import RuntimePaths
 28from .protocol import send_request
 29from .url_normalization import normalize_url
 30
 31EXIT_OK = 0
 32EXIT_ERROR = 1
 33_START_INSTRUCTION = "Start the daemon with: agent-search-gateway start"
 34
 35SocketClient = Callable[[Path, Request], Awaitable[Response]]
 36
 37
 38class DaemonLike(Protocol):
 39    async def start(self) -> None: ...
 40
 41
 42DaemonFactory = Callable[..., DaemonLike]
 43LoggingConfigurer = Callable[..., DebugLoggingSession]
 44DoctorRunner = Callable[..., Awaitable[DoctorReport]]
 45
 46
 47def build_parser() -> argparse.ArgumentParser:
 48    parser = argparse.ArgumentParser(prog="agent-search-gateway")
 49    subparsers = parser.add_subparsers(dest="command", required=True)
 50    start = subparsers.add_parser("start")
 51    start.add_argument("--debug", action="store_true")
 52    subparsers.add_parser("stop")
 53    subparsers.add_parser("doctor")
 54
 55    keyword = subparsers.add_parser("keyword-search")
 56    keyword.add_argument("query")
 57
 58    paper = subparsers.add_parser("paper-search")
 59    paper.add_argument("query")
 60
 61    llm = subparsers.add_parser("llm-search")
 62    llm.add_argument("prompt")
 63    llm.add_argument("--scope", choices=("web", "paper", "all"), default="web")
 64
 65    fetch = subparsers.add_parser("url-fetch")
 66    fetch.add_argument("url")
 67    fetch.add_argument("focus", nargs="?")
 68    return parser
 69
 70
 71def _request_from_args(args: argparse.Namespace) -> Request:
 72    if args.command == "stop":
 73        return ShutdownRequest()
 74    if args.command == "keyword-search":
 75        query = args.query.strip()
 76        if not query:
 77            raise InputFailure(ErrorCode.EMPTY_QUERY, "Query must not be empty")
 78        return KeywordSearchRequest(query)
 79    if args.command == "paper-search":
 80        query = args.query.strip()
 81        if not query:
 82            raise InputFailure(ErrorCode.EMPTY_QUERY, "Query must not be empty")
 83        return PaperSearchRequest(query)
 84    if args.command == "llm-search":
 85        prompt = args.prompt.strip()
 86        if not prompt:
 87            raise InputFailure(ErrorCode.EMPTY_QUERY, "Prompt must not be empty")
 88        return LLMSearchRequest(prompt, args.scope)
 89    if args.command == "url-fetch":
 90        url = normalize_url(args.url)
 91        focus = args.focus.strip() if args.focus is not None and args.focus.strip() else None
 92        return URLFetchRequest(str(url), focus)
 93    raise InputFailure(ErrorCode.BAD_REQUEST, "Unknown command")
 94
 95
 96def _write_text(stream: TextIO, text: str) -> None:
 97    stream.write(text)
 98    if not text.endswith("\n"):
 99        stream.write("\n")
100
101
102async def run_command(
103    args: argparse.Namespace,
104    paths: RuntimePaths,
105    *,
106    client: SocketClient = send_request,
107    daemon_factory: DaemonFactory = ForegroundDaemon,
108    logging_configurer: LoggingConfigurer = configure_debug_logging,
109    doctor_runner: DoctorRunner = run_doctor,
110    environ: Mapping[str, str] | None = None,
111    stdout: TextIO,
112    stderr: TextIO,
113) -> int:
114    if args.command == "start":
115        loop = asyncio.get_running_loop()
116        current_task = asyncio.current_task()
117        terminated = False
118        signal_handler_installed = False
119        logging_session: DebugLoggingSession | None = None
120
121        def cancel_for_sigterm() -> None:
122            nonlocal terminated
123            terminated = True
124            if current_task is not None:
125                current_task.cancel()
126
127        try:
128            if args.debug:
129                logging_session = logging_configurer(paths.debug_log_file, stderr=stderr)
130            daemon = daemon_factory(
131                paths,
132                debug=args.debug,
133                logging_session=logging_session,
134            )
135            try:
136                loop.add_signal_handler(signal.SIGTERM, cancel_for_sigterm)
137                signal_handler_installed = True
138            except (NotImplementedError, RuntimeError, ValueError):
139                pass
140            await daemon.start()
141        except asyncio.CancelledError:
142            if terminated:
143                return 128 + signal.SIGTERM
144            raise
145        except GatewayError as exc:
146            _write_text(stderr, exc.message)
147            return EXIT_ERROR
148        finally:
149            if signal_handler_installed:
150                loop.remove_signal_handler(signal.SIGTERM)
151            if logging_session is not None:
152                logging_session.close()
153        return EXIT_OK
154
155    if args.command == "doctor":
156        try:
157            report = await doctor_runner(
158                paths,
159                environ=os.environ if environ is None else environ,
160            )
161        except Exception:
162            _write_text(stderr, "[fail] doctor internal error")
163            return EXIT_ERROR
164        render_doctor(report, stdout)
165        return report.exit_code
166
167    try:
168        request = _request_from_args(args)
169    except GatewayError as exc:
170        _write_text(stderr, exc.message)
171        return EXIT_ERROR
172
173    try:
174        response = await client(paths.socket_file, request)
175    except DaemonUnavailable:
176        if isinstance(request, ShutdownRequest):
177            _write_text(stdout, "Daemon is not running.")
178            return EXIT_OK
179        _write_text(stderr, _START_INSTRUCTION)
180        return EXIT_ERROR
181    except GatewayError as exc:
182        _write_text(stderr, exc.message)
183        return EXIT_ERROR
184
185    if isinstance(response, SuccessResponse):
186        _write_text(stdout, response.text)
187        return EXIT_OK
188    if isinstance(response, ErrorResponse):
189        _write_text(stderr, response.message)
190        return EXIT_ERROR
191    _write_text(stderr, "Invalid daemon response")
192    return EXIT_ERROR
193
194
195def main(argv: list[str] | None = None) -> int:
196    parser = build_parser()
197    args = parser.parse_args(argv)
198    return asyncio.run(
199        run_command(
200            args,
201            RuntimePaths.default(),
202            stdout=sys.stdout,
203            stderr=sys.stderr,
204        )
205    )
EXIT_OK = 0
EXIT_ERROR = 1
class DaemonLike(typing.Protocol):
39class DaemonLike(Protocol):
40    async def start(self) -> None: ...

Base class for protocol classes.

Protocol classes are defined as::

class Proto(Protocol):
    def meth(self) -> int:
        ...

Such classes are primarily used with static type checkers that recognize structural subtyping (static duck-typing).

For example::

class C:
    def meth(self) -> int:
        return 0

def func(x: Proto) -> int:
    return x.meth()

func(C())  # Passes static type check

See PEP 544 for details. Protocol classes decorated with @typing.runtime_checkable act as simple-minded runtime protocols that check only the presence of given attributes, ignoring their type signatures. Protocol classes can be generic, they are defined as::

class GenProto[T](Protocol):
    def meth(self) -> T:
        ...
DaemonLike(*args, **kwargs)
1739def _no_init_or_replace_init(self, *args, **kwargs):
1740    cls = type(self)
1741
1742    if cls._is_protocol:
1743        raise TypeError('Protocols cannot be instantiated')
1744
1745    # Already using a custom `__init__`. No need to calculate correct
1746    # `__init__` to call. This can lead to RecursionError. See bpo-45121.
1747    if cls.__init__ is not _no_init_or_replace_init:
1748        return
1749
1750    # Initially, `__init__` of a protocol subclass is set to `_no_init_or_replace_init`.
1751    # The first instantiation of the subclass will call `_no_init_or_replace_init` which
1752    # searches for a proper new `__init__` in the MRO. The new `__init__`
1753    # replaces the subclass' old `__init__` (ie `_no_init_or_replace_init`). Subsequent
1754    # instantiation of the protocol subclass will thus use the new
1755    # `__init__` and no longer call `_no_init_or_replace_init`.
1756    for base in cls.__mro__:
1757        init = base.__dict__.get('__init__', _no_init_or_replace_init)
1758        if init is not _no_init_or_replace_init:
1759            cls.__init__ = init
1760            break
1761    else:
1762        # should not happen
1763        cls.__init__ = object.__init__
1764
1765    cls.__init__(self, *args, **kwargs)
async def start(self) -> None:
40    async def start(self) -> None: ...
DaemonFactory = collections.abc.Callable[..., DaemonLike]
LoggingConfigurer = collections.abc.Callable[..., agent_search_gateway.observability.DebugLoggingSession]
DoctorRunner = collections.abc.Callable[..., collections.abc.Awaitable[agent_search_gateway.doctor.DoctorReport]]
def build_parser() -> argparse.ArgumentParser:
48def build_parser() -> argparse.ArgumentParser:
49    parser = argparse.ArgumentParser(prog="agent-search-gateway")
50    subparsers = parser.add_subparsers(dest="command", required=True)
51    start = subparsers.add_parser("start")
52    start.add_argument("--debug", action="store_true")
53    subparsers.add_parser("stop")
54    subparsers.add_parser("doctor")
55
56    keyword = subparsers.add_parser("keyword-search")
57    keyword.add_argument("query")
58
59    paper = subparsers.add_parser("paper-search")
60    paper.add_argument("query")
61
62    llm = subparsers.add_parser("llm-search")
63    llm.add_argument("prompt")
64    llm.add_argument("--scope", choices=("web", "paper", "all"), default="web")
65
66    fetch = subparsers.add_parser("url-fetch")
67    fetch.add_argument("url")
68    fetch.add_argument("focus", nargs="?")
69    return parser
async def run_command( args: argparse.Namespace, paths: agent_search_gateway.paths.RuntimePaths, *, client: Callable[[pathlib.Path, agent_search_gateway.models.KeywordSearchRequest | agent_search_gateway.models.PaperSearchRequest | agent_search_gateway.models.LLMSearchRequest | agent_search_gateway.models.URLFetchRequest | agent_search_gateway.models.ShutdownRequest], Awaitable[agent_search_gateway.models.SuccessResponse | agent_search_gateway.models.ErrorResponse]] = <function send_request>, daemon_factory: Callable[..., DaemonLike] = <class 'agent_search_gateway.daemon.ForegroundDaemon'>, logging_configurer: Callable[..., agent_search_gateway.observability.DebugLoggingSession] = <function configure_debug_logging>, doctor_runner: Callable[..., Awaitable[agent_search_gateway.doctor.DoctorReport]] = <function run_doctor>, environ: Mapping[str, str] | None = None, stdout: <class 'TextIO'>, stderr: <class 'TextIO'>) -> int:
103async def run_command(
104    args: argparse.Namespace,
105    paths: RuntimePaths,
106    *,
107    client: SocketClient = send_request,
108    daemon_factory: DaemonFactory = ForegroundDaemon,
109    logging_configurer: LoggingConfigurer = configure_debug_logging,
110    doctor_runner: DoctorRunner = run_doctor,
111    environ: Mapping[str, str] | None = None,
112    stdout: TextIO,
113    stderr: TextIO,
114) -> int:
115    if args.command == "start":
116        loop = asyncio.get_running_loop()
117        current_task = asyncio.current_task()
118        terminated = False
119        signal_handler_installed = False
120        logging_session: DebugLoggingSession | None = None
121
122        def cancel_for_sigterm() -> None:
123            nonlocal terminated
124            terminated = True
125            if current_task is not None:
126                current_task.cancel()
127
128        try:
129            if args.debug:
130                logging_session = logging_configurer(paths.debug_log_file, stderr=stderr)
131            daemon = daemon_factory(
132                paths,
133                debug=args.debug,
134                logging_session=logging_session,
135            )
136            try:
137                loop.add_signal_handler(signal.SIGTERM, cancel_for_sigterm)
138                signal_handler_installed = True
139            except (NotImplementedError, RuntimeError, ValueError):
140                pass
141            await daemon.start()
142        except asyncio.CancelledError:
143            if terminated:
144                return 128 + signal.SIGTERM
145            raise
146        except GatewayError as exc:
147            _write_text(stderr, exc.message)
148            return EXIT_ERROR
149        finally:
150            if signal_handler_installed:
151                loop.remove_signal_handler(signal.SIGTERM)
152            if logging_session is not None:
153                logging_session.close()
154        return EXIT_OK
155
156    if args.command == "doctor":
157        try:
158            report = await doctor_runner(
159                paths,
160                environ=os.environ if environ is None else environ,
161            )
162        except Exception:
163            _write_text(stderr, "[fail] doctor internal error")
164            return EXIT_ERROR
165        render_doctor(report, stdout)
166        return report.exit_code
167
168    try:
169        request = _request_from_args(args)
170    except GatewayError as exc:
171        _write_text(stderr, exc.message)
172        return EXIT_ERROR
173
174    try:
175        response = await client(paths.socket_file, request)
176    except DaemonUnavailable:
177        if isinstance(request, ShutdownRequest):
178            _write_text(stdout, "Daemon is not running.")
179            return EXIT_OK
180        _write_text(stderr, _START_INSTRUCTION)
181        return EXIT_ERROR
182    except GatewayError as exc:
183        _write_text(stderr, exc.message)
184        return EXIT_ERROR
185
186    if isinstance(response, SuccessResponse):
187        _write_text(stdout, response.text)
188        return EXIT_OK
189    if isinstance(response, ErrorResponse):
190        _write_text(stderr, response.message)
191        return EXIT_ERROR
192    _write_text(stderr, "Invalid daemon response")
193    return EXIT_ERROR
def main(argv: list[str] | None = None) -> int:
196def main(argv: list[str] | None = None) -> int:
197    parser = build_parser()
198    args = parser.parse_args(argv)
199    return asyncio.run(
200        run_command(
201            args,
202            RuntimePaths.default(),
203            stdout=sys.stdout,
204            stderr=sys.stderr,
205        )
206    )