diff --git a/.github/workflows/contract.yml b/.github/workflows/contract.yml index 04eb4dc..8d4bfe3 100644 --- a/.github/workflows/contract.yml +++ b/.github/workflows/contract.yml @@ -18,14 +18,16 @@ jobs: python-version: '3.14' - run: uv sync --all-packages # SDK features land on development before the platform deploys to - # prod, so PRs targeting development check against the dev spec. + # prod, so every PR but a release into main checks against the dev + # spec — including one stacked on another feature branch, which is + # just as far ahead of prod as the branch it targets. # DEV_SPEC_URL is a repo secret; fork PRs don't receive it and # fall back to the prod spec, keeping dev infra internal-only. - env: BASE_REF: ${{ github.base_ref }} DEV_SPEC_URL: ${{ secrets.DEV_SPEC_URL }} run: | - if [ "$BASE_REF" = "development" ] && [ -n "$DEV_SPEC_URL" ]; then + if [ "$BASE_REF" != "main" ] && [ -n "$DEV_SPEC_URL" ]; then uv run python scripts/check_contract.py --spec-url "$DEV_SPEC_URL" uv run python scripts/gen_requests.py --check --spec-url "$DEV_SPEC_URL" else diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml index 5717236..c29f46e 100644 --- a/.pre-commit-config.yaml +++ b/.pre-commit-config.yaml @@ -39,13 +39,13 @@ repos: hooks: - id: foxguard name: foxguard - entry: npx -y foxguard@0.10.0 --staged + entry: npx -y foxguard@0.14.0 --staged language: system pass_filenames: false always_run: true - id: foxguard-secrets name: foxguard-secrets - entry: npx -y foxguard@0.10.0 secrets --staged + entry: npx -y foxguard@0.14.0 secrets --staged language: system pass_filenames: false always_run: true diff --git a/CHANGELOG.md b/CHANGELOG.md index 4bae8c0..9283b6b 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,5 +1,36 @@ # Changelog +## 0.4.0 (2026-09-23) + +- SDK: `validate_icp` can run without an LLM key. Pass `integration_id=NATIVE_ICP_ENGINE` (exported from `discolike`, the string `"native-icp"`) to score with DiscoLike's own ICP-fit model instead of your BYOK LLM: no LLM cost, no LLM key, and no web search on that run. The same sentinel works on `discogen.process` for a prompt that already carries the validation structure. Two new errors are specific to it: a 400 `ValidationError` when the ICP text does not yield a Mandatory / Reject if / Nice-to-have prompt, and a 503 `ServerError` when no ICP-fit engine is available. Task lifecycle, polling and statuses are unchanged. +- SDK: `Job` / `AsyncJob` returned by `validate_icp`, `discogen.process` and `discogen.process_personas` carry `column_name` from the submit response. The engine picks the columns — an LLM validation returns `Fit` / `Confidence` / `Reasoning`, the native model `ICP Fit` (`Yes` / `No` at a 0.50 threshold on the score) / `ICP Score` (the calibrated probability, 0.00-1.00 as a string) / `Reasoning (always null)` — so read it rather than hardcoding either set. It is `None` on a job reattached with `discogen.job(task_id)`. +- CLI: `validate-icp --integration-id native-icp` and `discogen run --integration-id native-icp` select the native ICP-fit model; `discolike validate-icp --help` documents both column sets. +- **Breaking:** the ICP-fit verdict is title case on every surface. An LLM `validate_icp` run's `Fit` column now returns `Yes` / `No` instead of `yes` / `no`, matching the native engine's `ICP Fit` column, DiscoGen and the app. Code that matches the verdict on the exact string `"yes"` has to be updated; compare case-insensitively to survive either. +- SDK/CLI: `discogen.process` and `discogen.process_personas` take `typed_columns`, and the CLI gains `--typed-columns`. It opts the query into the detector answering yes/no, fixed-set and scale columns with a TypeSafe judgment model instead of generated prose; which columns are typed and how is decided server-side from the query text, not passed in by the caller. Off by default. See `include_confidence` below for the sibling flag that adds a confidence column to each typed column. +- SDK/CLI: `discogen.process` and `discogen.process_personas` take `include_confidence`, and the CLI gains `--include-confidence`. It applies to typed columns, which a DiscoGen run answers with a TypeSafe judgment model rather than generated prose; with it on, each typed column gets a sibling confidence column. Off by default, and it changes display only - the answers and their probabilities come back either way. +- SDK: `contacts.generate` can run without an LLM key. Pass `integration_id=NATIVE_ENGINE` (exported from `discolike`) for DiscoLike Groove, DiscoLike's native extractor, or omit `integration_id` to get your default LLM integration and native when you have none. The native engine returns every person the search surfaces with a title and does not validate titles against `icp_text`; a search provider is still required. `JobStatus` gains `title_validation` (`"llm"` or `"none"`, `None` on non-ContaGen jobs) so you can tell which engine ran. +- SDK: discover and count take `sub_industry` / `negate_sub_industry` (217 second-level labels scoped to their parent industry category; a bare label like `ROOFING` adds `CONSTRUCTION` to `category` server-side) and a multi-shape geo filter: `geo` (`lat,lon` or `lat,lon,radius`), `bbox` (`min_lat,min_lon,max_lat,max_lon`, longitudes may wrap the antimeridian) and `lat` / `lon` / `radius` (`50km`, `30mi`, or a bare number meaning kilometres; defaults to 50km). `geo` and `bbox` are lists, every circle and box is OR'd with the others and with the `lat`/`lon` centre, and a query carries at most 10 shapes. Passing a single `bbox` or `geo` string still works and is sent as a one-element list. The SDK applies the platform's own shape rules before the request leaves the client: a box is four finite numbers with `-90 <= min_lat < max_lat <= 90`, longitudes within ±180 and `min_lon` different from `max_lon`; a circle is a valid coordinate pair with an optional radius above 0 and at most 1000km. The cross-field rules are checked locally too: `lat` and `lon` come together, `radius` needs them and follows the same radius rule, and the centre, `geo` and `bbox` total at most 10 shapes. +- SDK: `CompanyProfile` gains `sub_industry` (label:confidence, `None` when the domain was never scored against its category's sub-labels), `latitude`, `longitude` and `geo_precision`. The coordinates come back under their full names, and a response still using the pre-rename `lat` / `lon` keys parses into them, so the SDK works against an API on either side of that release. The `lat` / `lon` discover and count filters keep their names. +- CLI: `discover` and `count` gain `--sub-industry`, `--negate-sub-industry`, `--lat`, `--lon`, `--radius`, `--geo` and `--bbox`. `--geo` and `--bbox` are repeatable, one flag per shape. +- CLI: error envelopes on stderr gain a stable snake_case `code` (`validation_error`, `auth_required`, `auth_invalid`, `plan_access`, `rate_limited`, `network_error`, `not_found`, `server_error`, `job_failed`, `job_timeout`) and an `exit_code`, next to the existing `error` class name, `message`, and `status_code`. Agents and scripts should branch on `code` or the exit code. +- CLI: `discolike --help` ends with the output contract (JSON on stdout, envelope on stderr, when JSON is the default), the exit-code table, environment variables, and jq examples. `discover`, `count`, `match`, `append`, `validate-icp`, `segment`, `extract`, `signup`, `contacts search`, `contacts generate`, `discogen run`, `discogen status`, `queries create-exclusion-list`, `auth login`, `auth status`, and `account usage` each document their success JSON shape and common error codes under "Output (success, exit 0)" and "Common errors". Epilogs print verbatim instead of being reflowed. +- CLI: the console entry point is now `discolike_cli.main:run`. Parser failures (unknown flag, bad typed value, missing argument, `typer.BadParameter`) print the same `{"code": "validation_error", ...}` envelope on stderr with exit 2 instead of click's usage text; `--help`, `--version`, and the bare `discolike` help screen are unchanged; Ctrl-C exits 130. +- CLI: `discolike auth login` writes its success JSON (`{"logged_in": true, ...}`) to stdout like every other command; the URL to open and other progress lines stay on stderr. +- CLI: `auth_required` is reported only when no credential was found at all; a stored OAuth credential that fails to refresh (expired session, malformed token response) reports `auth_invalid`. +- CLI: `discolike --version` prints the bare CLI version when stdout is not a TTY; the decorated line with the SDK version is unchanged on a terminal. +- CLI: domain lists from files — `--domains-file PATH` (CSV with a `domain` column, or one domain per line; merged with `--domain`) on `queries create-exclusion-list`, `contacts discover`, `contacts search`, `contacts count`, `contacts generate` and `discogen run`; `validate-icp --file` gains the `--domains-file` alias and reads the same shapes. `--domain` on `contacts generate` / `discogen run` is now optional when the file is given. +- CLI: `--params-file PATH` on `discover`, `count`, `contacts discover|search|count` — a JSON object of API parameter names (an app form copied over), validated locally with the SDK request model before the call. Precedence: file < `--param` < flags. +- CLI: `discover --exclude-domains-file PATH` merges a file into the inline `exclude_domain` list (100 max, checked locally). +- CLI: `discolike bulk companies|estimate|contacts` — the volume pipeline as plain CLI calls. `companies` loops `discover` at up to 10,000 per page, saves each page as an exclusion list (`-round-N`) for the next page, appends to `--out` as pages land and resumes from that CSV on rerun, re-excluding every domain already in it even when `--exclusion-query-id` is given (warns at the 250,000-domain exclusion capacity). `estimate` sums the free `contacts count` over 1,000-domain slices and reports the upper bound at `--per-company`. `contacts` slices the domain list at `10000 / per-company` per `contacts discover` call (`results_by_company`, no `offset`), flattens to one row per contact, checkpoints finished slices in `.checkpoint` (stamped with the domains, `--per-company` and filters it was written for, so a rerun with different inputs is refused instead of silently skipped) and skips them on rerun. All three validate the request locally before the first billable call, keep one call in flight under `--rate-limit`, retry on 429/5xx, and print a JSON summary on stdout. +- SDK: `ContactGenerateRequest.find_emails` (default `False`) asks the platform to run the email finder over every named, email-less row before the ContaGen job completes, filling `email` and `email_status` on those rows. Found addresses bill under the finder's rules; the rest of the job stays unbilled. +- CLI: `discolike contacts generate --find-emails` sets the same flag, so a generate run can return name + title + email triples without chaining `email find-batch` afterwards. +- SDK: OAuth refreshes now resend the RFC 8707 `resource` the token was issued for. `OAuthCredential` gains an optional `resource` field, filled in by `exchange_code` and persisted to the config file; credentials stored by earlier releases load with `resource=None` and refresh as before until the next `discolike auth login`. Without it, an authorization server configured with a default resource could re-bind a refreshed REST token to another audience. +- SDK (behavior change, no code change): company `address.state` now comes back from the API as the subdivision name ("California", "Tokyo") instead of the ISO code ("CA", "13"). `CompanyAddress.state` is still `str | None` and needs no migration, but anything joining or grouping on that value as a code has to resolve it. The contact's own `state` is unchanged and stays a code. +- SDK: state filters accept a code or a name, resolved server-side against the countries you selected. Discover/count still take one `country` value, but that value may be a region alias (`EU`, `APAC`, `DACH`) and the state resolves against every member. Contacts state filters accept multiple countries and drop a value they cannot resolve rather than erroring. +- SDK: `MatchCompanyParams.state` works for any country with subdivisions, not just the US, and takes a code or a name. +- SDK: regenerated request models — the state field descriptions above now ship in `discolike.requests`. +- SDK (note for maintainers): `geo` and the widened `bbox` were written into `discolike/_generated/requests.py` by hand, ahead of the platform deploy that puts them in the OpenAPI spec. `scripts/gen_requests.py` rebuilds that file wholesale from the spec, so until the deploy lands `--check` reports those fields as a diff; that is the pending release, not drift. Regenerate once api-dev serves them. + ## 0.3.2 (2026-09-02) - CLI: every SDK request field now has a flag — `discover`/`count` gain `--variance`, `--min-similarity`, `--consensus`, `--inclusion-query-id`, `--language`, `--social`, `--subdomain`, `--start-date`, `--redirect`, `--exclude-leadgen` and the `--auto-*` toggles; contacts `search`/`count`/`discover` gain the full filter set; `match` gains per-column flags for file mode and `--min-match-confidence`; `append`/`segment` take `--query-id`; `extract` accepts `--domain`. Dict-typed fields stay `--param` only. diff --git a/README.md b/README.md index 5a81ba4..37ff589 100644 --- a/README.md +++ b/README.md @@ -137,6 +137,45 @@ for company in companies: print(company.domain, company.name, company.similarity) ``` +Narrow to a sub-industry within a radius of a point: + +```python +from discolike import Discolike +from discolike.requests import DiscoverParams + +client = Discolike() + +roofers = client.discover( + DiscoverParams( + sub_industry=["ROOFING"], + lat=30.2672, + lon=-97.7431, + radius="50mi", + max_records=25, + ) +) +for company in roofers: + print(company.domain, company.name, company.sub_industry) +``` + +A bare `ROOFING` resolves to `CONSTRUCTION/ROOFING` and adds `CONSTRUCTION` to the category filter +server-side; pass the qualified form yourself if you would rather be explicit. + +One query can cover several areas at once. `geo` takes `lat,lon` or `lat,lon,radius` and `bbox` takes +`min_lat,min_lon,max_lat,max_lon`; both are lists, and every circle, every box and the `lat`/`lon` +centre are OR'd together, up to 10 shapes: + +```python +austin_dallas_and_houston = client.discover( + DiscoverParams( + sub_industry=["ROOFING"], + geo=["30.2672,-97.7431,30mi", "32.7767,-96.797,30mi"], + bbox=["29.6,-95.7,30.1,-95.0"], + max_records=25, + ) +) +``` + Run DiscoGen research over a set of domains and wait for the result: ```python @@ -270,6 +309,50 @@ result = job.wait() `JobTimeoutError` is a client-side wait limit only — the task keeps running server-side (large DiscoGen runs can take hours), so call `wait()` again to resume or fetch `status()` later. Cancelled tasks still return results for every item that finished before cancellation. Send one job per list (up to 10,000 domains) rather than splitting into parallel jobs — concurrent DiscoGen jobs share your LLM provider key and slow each other down. +### Contact generation without an LLM key + +`contacts.generate` runs on your own search provider plus either your own LLM or DiscoLike Groove, DiscoLike's native extractor. Pass `NATIVE_ENGINE` to skip the LLM entirely: + +```python +from discolike import NATIVE_ENGINE +from discolike.requests import ContactGenerateRequest + +job = client.contacts.generate( + ContactGenerateRequest( + icp_text="VPs or Directors of Marketing at B2B SaaS", + domains=["gusto.com", "rippling.com"], + integration_id=NATIVE_ENGINE, + ) +) +result = job.wait() +print(result.title_validation) # "none" on the native engine, "llm" on a BYOK run +``` + +The native engine returns every person the search surfaces with a title and does not validate titles against `icp_text`, so filter them yourself when that matters. Omit `integration_id` to use your default LLM integration, or native when you have none. A search provider is required either way. + +### ICP validation without an LLM key + +`validate_icp` runs your ICP text against each domain on your own LLM provider key, or on DiscoLike's own ICP-fit model. Pass `NATIVE_ICP_ENGINE` for the latter — no LLM key, no LLM cost, and no web search on that run: + +```python +from discolike import NATIVE_ICP_ENGINE +from discolike.requests import ValidateIcpRequest + +job = client.validate_icp( + ValidateIcpRequest( + icp_text="Cybersecurity for SMBs in North America, 50-500 employees", + domains=["gusto.com", "rippling.com"], + integration_id=NATIVE_ICP_ENGINE, + ) +) +print(job.column_name) # ["ICP Fit", "ICP Score", "Reasoning"] +result = job.wait() +``` + +The engine decides the result columns, so read `job.column_name` instead of hardcoding them: an LLM run returns `Fit` / `Confidence` / `Reasoning`, the native model `ICP Fit` / `ICP Score` / `Reasoning (always null)`. `ICP Fit` is `Yes` or `No` at a 0.50 threshold on `ICP Score`, the calibrated probability as a 0.00-1.00 string. The native model returns only that score, so `Reasoning` is always `null` — the column is there to keep the set the same shape as an LLM run, not to carry an explanation. `integration_id="native-icp"` also works on `discogen.process` for a prompt that already carries the validation structure. + +Two errors are specific to the native engine: a 400 `ValidationError` when the ICP text does not yield a Mandatory / Reject if / Nice-to-have prompt, and a 503 `ServerError` when no ICP-fit engine is available. Task lifecycle, polling and statuses are the same either way. + ### Error handling All errors inherit from `DiscolikeError`: diff --git a/examples/README.md b/examples/README.md index d503c24..45d69b9 100644 --- a/examples/README.md +++ b/examples/README.md @@ -18,6 +18,6 @@ export DISCOLIKE_API_KEY="dl_..." # create one at https://app.discolike.com/ac | [`discover_and_enrich.py`](discover_and_enrich.py) | Discovers companies for an ICP, then runs a DiscoGen research prompt over them (needs a BYOK LLM provider) | `python examples/discover_and_enrich.py --icp "Cybersecurity for SMBs" --country US --query "What is their pricing model?"` | | [`find_emails_from_csv.py`](find_emails_from_csv.py) | Finds verified work emails for a CSV of first name, last name, domain in batches of 500; only status `found` bills | `python examples/find_emails_from_csv.py people.csv --output emails.csv` | | [`match_crm_contacts.py`](match_crm_contacts.py) | Matches a messy CRM contact export to DiscoLike persona IDs with resumable checkpointing | `python examples/match_crm_contacts.py contacts.csv --output matched.csv` | -| [`cli_recipes.sh`](cli_recipes.sh) | The same searches as `discolike discover`, `discolike count`, `discolike contacts search`, and `discolike signup` one-liners | `bash examples/cli_recipes.sh` | +| [`cli_recipes.sh`](cli_recipes.sh) | The same searches as `discolike discover`, `discolike count`, `discolike contacts search`, and `discolike signup` one-liners, plus a `discolike bulk` volume pull | `bash examples/cli_recipes.sh` | Every script prints `--help`. Employee ranges are `min,max` strings such as `51,200`; countries are ISO-2 codes or region aliases like `EU`, `DACH`, `APAC`. diff --git a/examples/cli_recipes.sh b/examples/cli_recipes.sh index f3e984a..4622818 100644 --- a/examples/cli_recipes.sh +++ b/examples/cli_recipes.sh @@ -27,3 +27,10 @@ discolike company data stripe.com --format json # agent_signup_to_first_search.py: open an account for a person, no auth needed discolike signup --email jane@acme.com --first-name Jane --last-name Doe --agent cookbook + +# Volume: every company that matches, then N contacts at each, checkpointed and resumable. +# form.json is a JSON object of API parameter names (an app Discover form copied over); count is free. +discolike count --params-file form.json --format json +discolike bulk companies --params-file form.json --max-companies 50000 --run-name agencies --out companies.csv +discolike bulk estimate --domains-file companies.csv --per-company 10 --summary "growth marketing lead" # free +discolike bulk contacts --domains-file companies.csv --per-company 10 --summary "growth marketing lead" --out contacts.csv diff --git a/packages/discolike-cli/README.md b/packages/discolike-cli/README.md index bab4b42..82f4262 100644 --- a/packages/discolike-cli/README.md +++ b/packages/discolike-cli/README.md @@ -35,12 +35,25 @@ discolike company data stripe.com discolike extract https://stripe.com/enterprise ``` -Top-level commands: `discover`, `count`, `match`, `extract`, `validate-icp`, `append`, `segment` — plus `auth`, `company`, `contacts`, `discogen`, `queries`, `account`, `search-providers`, and `llm-providers` command groups. +Top-level commands: `discover`, `count`, `match`, `extract`, `validate-icp`, `append`, `segment` — plus `auth`, `bulk`, `company`, `contacts`, `discogen`, `queries`, `account`, `search-providers`, and `llm-providers` command groups. + +### Volume pulls + +`discolike bulk` walks past the 10,000-per-search ceiling and pulls contacts for a whole domain list into one CSV, with checkpoint and resume: + +```bash +discolike bulk companies --params-file form.json --max-companies 50000 --run-name agencies --out companies.csv +discolike bulk estimate --domains-file companies.csv --per-company 10 # free size check +discolike bulk contacts --domains-file companies.csv --per-company 10 --summary "growth marketing" --out contacts.csv +``` + +`companies` saves each page as an exclusion list (`-round-N`) and excludes it from the next page; rerunning with the same `--out` resumes from the CSV. `contacts` slices the domain list at `10000 / per-company` domains per call and records finished slices in `.checkpoint`. Both keep one call in flight under `--rate-limit` (default 10/min, the Pro rate on `/discover` and `/contacts`), retry on 429/5xx, and print a JSON summary at the end. Filters come from `--params-file`, `--param` and the common flags; the paging fields are managed for you. ### Conventions - Results print as JSON to stdout; errors print as JSON (`error`, `message`, `status_code`) to stderr. - Pass `--format table` for a human-readable table — used automatically when stdout is a TTY. +- Volume inputs come from files: `--domains-file companies.csv` (a `domain` column, or one domain per line) on `queries create-exclusion-list`, `contacts discover|search|count|generate`, `discogen run` and `validate-icp`; `--params-file form.json` (a JSON object of API parameter names, e.g. an app form) on `discover`, `count` and `contacts discover|search|count`; `--exclude-domains-file` on `discover`. Precedence: file < `--param` < flags. - Async endpoints (`match --file`, `discogen run`, `discogen run-personas`, `segment`, `validate-icp`) take `--wait` to block until the job finishes. Without it, you get a `task_id` back to poll with `discolike discogen status --family `. `append` is synchronous — it returns enriched rows directly (or writes CSV bytes to `--output`). ### Exit codes diff --git a/packages/discolike-cli/pyproject.toml b/packages/discolike-cli/pyproject.toml index 911943f..656b253 100644 --- a/packages/discolike-cli/pyproject.toml +++ b/packages/discolike-cli/pyproject.toml @@ -4,7 +4,7 @@ build-backend = "hatchling.build" [project] name = "discolike-cli" -version = "0.3.2" +version = "0.4.0" description = "Official CLI for the DiscoLike API" readme = "README.md" license = "MIT" @@ -12,7 +12,7 @@ license-files = ["LICENSE"] requires-python = ">=3.10" authors = [{ name = "DiscoLike", email = "support@discolike.com" }] dependencies = [ - "discolike==0.3.2", + "discolike==0.4.0", "typer>=0.12", "rich>=13.0", ] @@ -34,7 +34,7 @@ classifiers = [ ] [project.scripts] -discolike = "discolike_cli.main:app" +discolike = "discolike_cli.main:run" [project.urls] Homepage = "https://www.discolike.com" diff --git a/packages/discolike-cli/src/discolike_cli/_help.py b/packages/discolike-cli/src/discolike_cli/_help.py new file mode 100644 index 0000000..7a9913e --- /dev/null +++ b/packages/discolike-cli/src/discolike_cli/_help.py @@ -0,0 +1,292 @@ +"""Help-text epilogs that spell out the CLI's output contract for agents and scripts. + +Every constant here is plain text (no Rich markup) so it renders identically in a +terminal, in a pipe, and in a test assertion. Response shapes mirror the SDK models in +``discolike.resources``; keep them in sync when a model changes. +""" + +from __future__ import annotations + +import textwrap + +import typer +from typer._click import Context +from typer._click import HelpFormatter +from typer.core import TyperCommand +from typer.core import TyperGroup + +from discolike._config import ENV_API_KEY +from discolike._config import config_path + +EPILOG_INDENT = " " + + +def _echo_verbatim(epilog: str) -> None: + # Typer's rich formatter reflows epilog paragraphs into prose, which destroys the + # aligned tables below. Print the epilog ourselves, exactly as written. + typer.echo(textwrap.indent(epilog.rstrip("\n"), EPILOG_INDENT)) + + +class ContractCommand(TyperCommand): + """A command whose epilog is printed verbatim, not reflowed.""" + + def format_help(self, ctx: Context, formatter: HelpFormatter) -> None: + epilog, self.epilog = self.epilog, None + try: + super().format_help(ctx, formatter) + finally: + self.epilog = epilog + if epilog: + _echo_verbatim(epilog) + + +class ContractGroup(TyperGroup): + """A command group whose epilog is printed verbatim, not reflowed.""" + + def format_help(self, ctx: Context, formatter: HelpFormatter) -> None: + epilog, self.epilog = self.epilog, None + try: + super().format_help(ctx, formatter) + finally: + self.epilog = epilog + if epilog: + _echo_verbatim(epilog) + + +_CONFIG_PATH_HINT = "$XDG_CONFIG_HOME/discolike/config.json (default ~/.config/discolike/config.json)" + +MAIN_EPILOG = f"""\ +Output contract: + Success JSON on stdout, exit 0. JSON is the default whenever stdout is not a + TTY; --format json forces it and --format table renders a table for + humans. Progress lines from --wait go to stderr only. + Failure {{"error", "code", "message", "status_code", "exit_code"}} on stderr + with a non-zero exit code. Branch on "code" (stable) or the exit code; + "error" is the SDK exception class name. rate_limited adds "retry_after". + Determinism No spinners, colors, or prompts on stdout when piped. + +Exit codes: + 0 success + 1 server_error, job_failed, job_timeout, or any other unrecoverable error + 2 validation_error input rejected before or by the API (bad flag, enum, range) + 3 auth_required no credential found: run `discolike auth login` + auth_invalid credential rejected (401/403); plan_access needs a higher plan + 4 rate_limited HTTP 429; wait "retry_after" seconds, then retry + 5 network_error could not reach the API + 6 not_found HTTP 404 + +Environment: + {ENV_API_KEY} API key; overrides the config file written by `discolike auth login`. + Config file {_CONFIG_PATH_HINT} + (resolved now: {config_path()}) + +Examples: + $ discolike count --country DE --phrase-match "managed detection" --format json | jq .count + $ discolike discover --icp-prompt "cybersecurity for SMBs" --max-records 100 --format json | jq -r '.[].domain' + $ discolike validate-icp --icp "sells to hospitals" --domain a.com --domain b.com --wait --format json | jq '.' +""" + +_COMPANY_ROW = """\ + {"domain": , "name": , "similarity": <0-1>, "score": , + "description": , "employees": , "revenue_range": , + "address": {"city", "state", "country"}, "industry_groups": {: }, + "business_model": {: }, "mx_provider": , ...}""" + +_JOB_SUBMITTED = """\ + Without --wait: + {"task_id": , "task_family": , "hint": "poll with: discolike discogen status ..."}""" + +_JOB_ERRORS = """\ +Common errors: + validation_error (exit 2) bad flag value, empty domain list, or an unknown --model. + auth_required (exit 3) no credential; run `discolike auth login`. + plan_access (exit 3) the plan does not include this job type. + job_failed (exit 1) the job finished with an error (only with --wait). + job_timeout (exit 1) --timeout elapsed before the job finished; poll with `discolike discogen status`.""" + +_SEARCH_ERRORS = """\ +Common errors: + validation_error (exit 2) unknown filter value, bad range, or more than 10 seed domains. + auth_required (exit 3) no credential; run `discolike auth login`. + plan_access (exit 3) the plan does not allow this filter or record count. + rate_limited (exit 4) HTTP 429; retry after "retry_after" seconds.""" + +COMMAND_EPILOGS: dict[str, str] = { + "discover": f"""\ +Output (success, exit 0): + [ +{_COMPANY_ROW} + ] + One row per matched company, ranked by similarity. Billed per 1,000 new records; + records seen in the last 90 days are free. Count first: `discolike count` is free. + +{_SEARCH_ERRORS} +""", + "count": f"""\ +Output (success, exit 0): + {{"count": }} + Free. Same filters as `discover`; run it before any large pull. + +{_SEARCH_ERRORS} +""", + "match": """\ +Output (success, exit 0): + Single (--name): + {"query": {"name", "country", "state", "city", "zip", "phones"}, + "matches": [{"domain": , "name": , "match_confidence": <0-100>, ...}]} + Bulk (--file) without --wait: + {"task_id": , "task_family": , "hint": } + Bulk with --wait: one row per input row; rows the backend could not process carry + "match_error": "search_failed", which is not the same as no match. + +Common errors: + validation_error (exit 2) neither --name nor --file given, or a named column is missing. + auth_required (exit 3) no credential; run `discolike auth login`. + job_timeout (exit 1) --timeout elapsed on a bulk match; poll with `discolike discogen status`. +""", + "append": """\ +Output (success, exit 0): + JSON (default): [{"domain": , ...}] one row per input row. + --csv --output FILE: {"written": , "bytes": } and the CSV is on disk. + Datasets: bizdata, domain_status, redirects, growth, vendors, subdomains. + Up to 10,000 rows per request; billed per new company record. + +Common errors: + validation_error (exit 2) missing --domain-column, unknown --dataset, or --csv without --output. + auth_required (exit 3) no credential; run `discolike auth login`. + plan_access (exit 3) a requested dataset is not on the plan. +""", + "validate-icp": f"""\ +Output (success, exit 0): +{_JOB_SUBMITTED} + With --wait: the results object keyed by domain: + {{: {{"Fit": "Yes"|"No", "Confidence": "high"|"medium"|"low", "Reasoning": }}, ...}} + Runs on the account's own LLM provider key. + With --integration-id native-icp: DiscoLike's own ICP-fit model, no LLM key or cost and no + web search. Keys become "ICP Fit" ("Yes"|"No" at a 0.50 threshold on the score), "ICP Score" + (the calibrated probability, 0.00-1.00 as a string) and "Reasoning", always null there: the + model returns a score, not an explanation. A 400 when the ICP text yields no Mandatory / + Reject if / Nice-to-have prompt, a 503 when no ICP-fit engine is available. + +{_JOB_ERRORS} +""", + "segment": f"""\ +Output (success, exit 0): +{_JOB_SUBMITTED} + With --wait: one BizData row per active domain, plus its cluster: + [{{"domain", "name", "score", "segment_id": , + "segment_description": , "probability": <0-1>, ...}}] + Closed or unindexed input domains are omitted. + +{_JOB_ERRORS} +""", + "extract": """\ +Output (success, exit 0): + {"text": , "language": } + Fetches the page live; use it to confirm a seed domain is the right company. + +Common errors: + validation_error (exit 2) neither a URL argument nor --domain given. + auth_required (exit 3) no credential; run `discolike auth login`. + not_found (exit 6) the page could not be fetched. +""", + "signup": """\ +Output (success, exit 0): + {"status": , "email": , "org_domain": , + "org_status": , "next_step": } + No credential is returned. Relay "next_step" to the person: they confirm by email, + log in, pick a plan, then run `discolike auth login`. No login is needed for this call. + +Common errors: + validation_error (exit 2) free-mail or disposable email domain, or a missing name. + error (exit 1) HTTP 409: the account already exists; send them to log in. +""", + "contacts search": """\ +Output (success, exit 0): + [{"persona_id": , "domain": , "name": , "title": , + "email": , ...}] + Billed per new contact record. `discolike contacts count` with the same filters is free. + +Common errors: + validation_error (exit 2) unknown seniority/department value or a bad range. + auth_required (exit 3) no credential; run `discolike auth login`. + plan_access (exit 3) contacts are not included on the plan. +""", + "contacts generate": f"""\ +Output (success, exit 0): +{_JOB_SUBMITTED} + With --wait: results keyed by domain, zero or more candidate rows each: + {{: [{{"name", "title", "department", "seniority", "email", "linkedin_url", + "skills", "phone": [{{"phone", "type"}}], "email_pattern", "email_pattern_confidence", + "email_pattern_guess"}}], ...}} + Rows are candidates; email is set only when publicly verifiable. + Runs live web search on the account's own LLM and search provider keys; no platform billing. + +{_JOB_ERRORS} +""", + "discogen run": f"""\ +Output (success, exit 0): +{_JOB_SUBMITTED} + With --wait: one entry per domain with the model's structured answer to --query. + Send the whole list in one call; parallel jobs share the provider key and rate-limit each other. + +{_JOB_ERRORS} +""", + "discogen status": """\ +Output (success, exit 0): + {"status": "pending"|"running"|"completed"|"failed", "progress": <0-100>, + "results": , "warnings": [], "estimated_cost": , + "cost_metadata": {: {...}, "search_provider": {...}}} + "results" is null until the job completes. Pass --family for non-DiscoGen jobs. + +Common errors: + validation_error (exit 2) unknown --family value. + auth_required (exit 3) no credential; run `discolike auth login`. + not_found (exit 6) no job with that task id in that family. +""", + "queries create-exclusion-list": """\ +Output (success, exit 0): + {"query_id": , "query_name": , "action": , + "domain_count": , "persona_id_count": , "row_count": , "tags": []} + Pass "query_id" as --exclusion-query-id (or --inclusion-query-id) on the next discover. + Minimum 20 domains; lists hold up to 250,000 domains on every plan. + +Common errors: + validation_error (exit 2) fewer than 20 domains, or a missing --name. + auth_required (exit 3) no credential; run `discolike auth login`. +""", + "auth login": """\ +Output (success, exit 0): + Browser (default): {"logged_in": true, "method": "oauth", "expires_at": } + --api-key KEY: {"logged_in": true, "source": "api_key"} + The credential is saved to the config file; later commands read it automatically. + Progress lines (the URL to open, browser wait) go to stderr; only the JSON above is on stdout. + Headless: --no-browser prints the URL to open; --port pins the loopback port for SSH. + +Common errors: + auth_invalid (exit 3) the API key was rejected. + error (exit 1) {"error": "LoginError", ...}: timeout, denied consent, or state mismatch. +""", + "auth status": """\ +Output (success, exit 0): + API key: {"source": "option"|"env"|"config", "method": "api_key", "api_key": , "valid": true} + OAuth: {"source": "config", "method": "oauth", "expires_at": , "expired": , "valid": true} + Exit 0 here is the whole proof that the CLI is signed in. + +Common errors: + auth_required (exit 3) no credential anywhere; run `discolike auth login`. + auth_invalid (exit 3) the stored credential no longer verifies; log in again. +""", + "account usage": """\ +Output (success, exit 0): + {"requests_mtd": , "records_mtd": , "spend_mtd": } + Month-to-date. Report it before any paid flow so the person knows the baseline. + +Common errors: + auth_required (exit 3) no credential; run `discolike auth login`. +""", +} + + +def epilog(name: str) -> str: + return COMMAND_EPILOGS[name] diff --git a/packages/discolike-cli/src/discolike_cli/_inputs.py b/packages/discolike-cli/src/discolike_cli/_inputs.py new file mode 100644 index 0000000..9728c88 --- /dev/null +++ b/packages/discolike-cli/src/discolike_cli/_inputs.py @@ -0,0 +1,64 @@ +"""File-backed inputs for volume flows: domain lists and whole-request JSON.""" + +from __future__ import annotations + +import csv +import json +import pathlib +from typing import Any + +import typer + +DOMAIN_COLUMN = "domain" +WWW_PREFIX = "www." + + +def read_domains_file(path: pathlib.Path) -> list[str]: + """CSV with a ``domain`` column, or one domain per line (first column when no header). + + Domains are stripped, lower-cased, ``www.``-less and de-duplicated in file order. + """ + try: + with path.open(newline="") as handle: + rows = list(csv.reader(handle)) + except FileNotFoundError as exc: + raise typer.BadParameter(f"domains file not found: {path}") from exc + except (OSError, UnicodeDecodeError) as exc: + raise typer.BadParameter(f"domains file {path} could not be read: {exc}") from exc + header = [cell.strip().lower() for cell in rows[0]] if rows else [] + column = header.index(DOMAIN_COLUMN) if DOMAIN_COLUMN in header else 0 + body = rows[1:] if DOMAIN_COLUMN in header else rows + seen: set[str] = set() + domains: list[str] = [] + for row in body: + if column >= len(row): + continue + domain = row[column].strip().lower().removeprefix(WWW_PREFIX) + if domain and domain not in seen: + seen.add(domain) + domains.append(domain) + if not domains: + raise typer.BadParameter(f"domains file has no domains: {path}") + return domains + + +def read_params_file(path: pathlib.Path) -> dict[str, Any]: + """JSON object of API parameter names, e.g. an app form copied over verbatim.""" + try: + loaded = json.loads(path.read_text()) + except FileNotFoundError as exc: + raise typer.BadParameter(f"params file not found: {path}") from exc + except (OSError, UnicodeDecodeError) as exc: + raise typer.BadParameter(f"params file {path} could not be read: {exc}") from exc + except json.JSONDecodeError as exc: + raise typer.BadParameter(f"params file {path} must contain valid JSON: {exc}") from exc + if not isinstance(loaded, dict): + raise typer.BadParameter(f"params file {path} must contain a JSON object") + return loaded + + +def merge_domains(inline: list[str] | None, file: pathlib.Path | None) -> list[str] | None: + """Inline ``--domain`` values first, then the file's, or ``None`` when neither was given.""" + if file is None: + return inline + return [*(inline or []), *read_domains_file(file)] diff --git a/packages/discolike-cli/src/discolike_cli/_output.py b/packages/discolike-cli/src/discolike_cli/_output.py index 593a8ed..0b3a7be 100644 --- a/packages/discolike-cli/src/discolike_cli/_output.py +++ b/packages/discolike-cli/src/discolike_cli/_output.py @@ -16,12 +16,16 @@ from rich.console import Console from rich.table import Table +from discolike._config import NO_CREDENTIAL_MESSAGE from discolike._exceptions import APIConnectionError from discolike._exceptions import AuthenticationError from discolike._exceptions import DiscolikeError +from discolike._exceptions import JobFailedError +from discolike._exceptions import JobTimeoutError from discolike._exceptions import NotFoundError from discolike._exceptions import PlanAccessError from discolike._exceptions import RateLimitError +from discolike._exceptions import ServerError from discolike._exceptions import ValidationError from discolike._models import DiscolikeModel @@ -40,6 +44,30 @@ } DEFAULT_EXIT_CODE = 1 +# Stable, snake_case error codes for agents and scripts to branch on. The class +# name in `error` is kept for backwards compatibility; `code` is the contract. +ERROR_CODES: dict[type, str] = { + ValidationError: "validation_error", + AuthenticationError: "auth_invalid", + PlanAccessError: "plan_access", + RateLimitError: "rate_limited", + APIConnectionError: "network_error", + NotFoundError: "not_found", + ServerError: "server_error", + JobFailedError: "job_failed", + JobTimeoutError: "job_timeout", +} +DEFAULT_ERROR_CODE = "error" +AUTH_REQUIRED_CODE = "auth_required" + + +def error_code(exc: DiscolikeError) -> str: + # The SDK raises the same exception class both when no credential exists and when a stored + # OAuth credential fails to refresh; only the former is "run `discolike auth login`" territory. + if isinstance(exc, AuthenticationError) and str(exc) == NO_CREDENTIAL_MESSAGE: + return AUTH_REQUIRED_CODE + return ERROR_CODES.get(type(exc), DEFAULT_ERROR_CODE) + class _JobStatusLike(Protocol): results: Any @@ -111,15 +139,18 @@ def emit(data: Any, *, fmt: str | None = None) -> None: # noqa: ANN401 -- accep def fail(exc: DiscolikeError) -> typer.Exit: + exit_code = EXIT_CODES.get(type(exc), DEFAULT_EXIT_CODE) payload: dict[str, Any] = { "error": type(exc).__name__, + "code": error_code(exc), "message": str(exc), "status_code": exc.status_code, + "exit_code": exit_code, } if isinstance(exc, RateLimitError) and exc.retry_after is not None: payload["retry_after"] = exc.retry_after print(json.dumps(payload), file=sys.stderr) - return typer.Exit(code=EXIT_CODES.get(type(exc), DEFAULT_EXIT_CODE)) + return typer.Exit(code=exit_code) def _accepts_list(model: type[pydantic.BaseModel], name: str) -> bool: diff --git a/packages/discolike-cli/src/discolike_cli/account.py b/packages/discolike-cli/src/discolike_cli/account.py index c09bd01..db0f204 100644 --- a/packages/discolike-cli/src/discolike_cli/account.py +++ b/packages/discolike-cli/src/discolike_cli/account.py @@ -2,6 +2,8 @@ import typer +from discolike_cli._help import ContractCommand +from discolike_cli._help import epilog from discolike_cli._output import emit from discolike_cli._output import handle_errors @@ -10,7 +12,7 @@ app = typer.Typer(help="Account usage and quota.") -@app.command("usage") +@app.command("usage", cls=ContractCommand, epilog=epilog("account usage")) @handle_errors def usage_command( ctx: typer.Context, diff --git a/packages/discolike-cli/src/discolike_cli/auth.py b/packages/discolike-cli/src/discolike_cli/auth.py index 9126a53..ccd86ba 100644 --- a/packages/discolike-cli/src/discolike_cli/auth.py +++ b/packages/discolike-cli/src/discolike_cli/auth.py @@ -34,6 +34,8 @@ from discolike._oauth import discover from discolike._oauth import exchange_code from discolike._oauth import register_client +from discolike_cli._help import ContractCommand +from discolike_cli._help import epilog from discolike_cli._loopback import CallbackServer from discolike_cli._output import emit from discolike_cli._output import handle_errors @@ -225,10 +227,10 @@ def _api_key_login(ctx: typer.Context, *, api_key: str | None) -> None: key = api_key or passed_globally or typer.prompt("API key", hide_input=True) _verify(ctx, api_key=key) save_config({"auth_method": AUTH_METHOD_API_KEY, "api_key": key}) - print(json.dumps({"logged_in": True, "source": AUTH_METHOD_API_KEY}), file=sys.stderr) + emit({"logged_in": True, "source": AUTH_METHOD_API_KEY}) -@app.command() +@app.command(cls=ContractCommand, epilog=epilog("auth login")) @handle_errors def login( ctx: typer.Context, @@ -263,13 +265,10 @@ def login( credential = _oauth_login(ctx, open_browser=not no_browser, port=port) _verify(ctx, auth=credential) save_credential(credential) - print( - json.dumps({"logged_in": True, "method": AUTH_METHOD_OAUTH, "expires_at": _iso(credential.expires_at)}), - file=sys.stderr, - ) + emit({"logged_in": True, "method": AUTH_METHOD_OAUTH, "expires_at": _iso(credential.expires_at)}) -@app.command() +@app.command(cls=ContractCommand, epilog=epilog("auth status")) @handle_errors def status(ctx: typer.Context) -> None: """Show which credential is in use (option, env, or config) and verify it against the API.""" diff --git a/packages/discolike-cli/src/discolike_cli/bulk.py b/packages/discolike-cli/src/discolike_cli/bulk.py new file mode 100644 index 0000000..7e622ee --- /dev/null +++ b/packages/discolike-cli/src/discolike_cli/bulk.py @@ -0,0 +1,561 @@ +"""``discolike bulk``: the volume pipeline — companies past 10k, contacts by domain slice, checkpoint and resume. + +Ports the mechanics of the support-side ``discolike_pipeline.py`` so an agent can drive the whole +flow with plain ``discolike`` calls (see issue #24). Progress goes to stderr, a JSON summary to stdout. +""" + +from __future__ import annotations + +import csv +import functools +import hashlib +import json +import pathlib +import sys +import time +from collections import deque +from collections.abc import Callable +from collections.abc import Iterator +from typing import Any +from typing import TypeVar + +import pydantic +import typer + +from discolike import Discolike +from discolike._exceptions import AuthenticationError +from discolike._exceptions import DiscolikeError +from discolike._exceptions import NotFoundError +from discolike._exceptions import PlanAccessError +from discolike._exceptions import RateLimitError +from discolike._exceptions import ValidationError +from discolike.requests import ContactFilters +from discolike.requests import ContactsCountParams +from discolike.requests import CreateExclusionListRequest +from discolike.requests import DiscoverParams +from discolike_cli._inputs import read_domains_file +from discolike_cli._output import build_request +from discolike_cli._output import emit +from discolike_cli._output import handle_errors +from discolike_cli.discover import DOMAINS_FILE_HELP +from discolike_cli.discover import PARAM_HELP +from discolike_cli.discover import PARAMS_FILE_HELP +from discolike_cli.discover import _merge_params + +T = TypeVar("T") + +MAX_RECORDS = 10000 # API ceiling per discover / contacts call +MIN_RECORDS = 20 # API floor on max_records +MAX_PER_COMPANY = 100 # API ceiling on results_by_company +MIN_EXCLUSION_LIST = 20 # saved lists reject anything smaller; shorter tails ride on exclude_domain +EXCLUSION_CAPACITY = 250_000 # domains an account can hold across exclusion lists (every plan) +COUNT_SLICE = 1000 # domains per free /contacts/count probe +RETRY_ATTEMPTS = 5 +RETRY_BASE_SECONDS = 5.0 +RETRY_MAX_SECONDS = 60.0 +CLIENT_TIMEOUT_SECONDS = 300.0 +DEFAULT_RATE_LIMIT_PER_MINUTE = 10 # Pro on /discover and /contacts; Starter 5, Team 15, Company 25, Enterprise 50 +COUNT_RATE_LIMIT_PER_MINUTE = 30 + +COMPANY_COLUMNS = ["domain", "name", "country", "employees", "similarity"] +CONTACT_COLUMNS = [ + "persona_id", "first_name", "last_name", "name", "title", "seniority", "department", + "email", "email_validated", "phone", "linkedin", "connections", + "domain", "company_name", "country", "state", "industry", "employees", "revenue_range", "jobstart_date", +] # fmt: skip +COMPANIES_MANAGED = ("exclusion_query_id", "exclude_domain", "max_records", "offset") +CONTACTS_MANAGED = ("domain", "max_records", "offset", "results_by_company", "max_companies") +FATAL = (ValidationError, AuthenticationError, PlanAccessError, NotFoundError, pydantic.ValidationError) + +RATE_LIMIT_HELP = "Calls per minute to stay under for the paid endpoint (Pro is 10 on /discover and /contacts)." +OVERWRITE_HELP = "Ignore an existing --out (and checkpoint) and start clean instead of resuming." +PER_COMPANY_HELP = "Contacts to pull per company (sets results_by_company and the domain slice size)." +EXCLUSION_QUERY_ID_HELP = "Saved query ID whose results are excluded (repeatable)." + +app = typer.Typer( + help=( + "Volume pulls with checkpoint and resume: walk the company index past the 10,000-per-search ceiling, " + "size a contact pull for free, then pull N contacts per company into one CSV." + ) +) + + +def _log(message: str) -> None: + print(message, file=sys.stderr, flush=True) + + +def _chunked(items: list[str], size: int) -> Iterator[tuple[int, list[str]]]: + for start in range(0, len(items), size): + yield start // size, items[start : start + size] + + +class _RateLimiter: + """Sliding-window limiter: at most ``per_minute`` acquisitions in any 60s window.""" + + def __init__(self, per_minute: int) -> None: + self._per_minute = per_minute + self._calls: deque[float] = deque() + + def acquire(self) -> None: + while True: + now = time.monotonic() + while self._calls and now - self._calls[0] > 60: + self._calls.popleft() + if len(self._calls) < self._per_minute: + self._calls.append(now) + return + time.sleep(60 - (now - self._calls[0]) + 0.05) + + +def _call_with_retry(limiter: _RateLimiter, fn: Callable[[], T]) -> T: + """The SDK transport already retries 429/5xx a few times; this layer survives longer outages.""" + for attempt in range(1, RETRY_ATTEMPTS + 1): + limiter.acquire() + try: + return fn() + except FATAL: + raise + except RateLimitError as err: + wait = min(RETRY_MAX_SECONDS, err.retry_after or RETRY_BASE_SECONDS * 2 ** (attempt - 1)) + _log(f" rate limited, retrying in {wait:.0f}s ({attempt}/{RETRY_ATTEMPTS})") + time.sleep(wait) + except DiscolikeError as err: + if attempt == RETRY_ATTEMPTS: + raise + _log(f" {type(err).__name__}: {err} - retry {attempt}/{RETRY_ATTEMPTS}") + time.sleep(RETRY_BASE_SECONDS * attempt) + raise RuntimeError("unreachable: retries exhausted without raising") + + +def _drop_managed(kwargs: dict[str, Any], managed: tuple[str, ...]) -> dict[str, Any]: + for key in managed: + if key in kwargs: + _log(f"note: {key} is managed by bulk, ignoring the supplied value") + kwargs.pop(key) + return kwargs + + +def _client(ctx: typer.Context) -> Discolike: + from discolike_cli.main import get_client + + return get_client(ctx).with_options(timeout=CLIENT_TIMEOUT_SECONDS) + + +# --------------------------------------------------------------------------- companies + + +@app.command("companies") +@handle_errors +def companies_command( + ctx: typer.Context, + icp_prompt: str | None = typer.Option(None, help="Natural-language ideal customer profile description."), + domain: list[str] | None = typer.Option(None, help="Seed domain for lookalike matching (repeatable)."), + phrase_match: list[str] | None = typer.Option(None, help="Phrase the company website must contain (repeatable)."), + negate_phrase_match: list[str] | None = typer.Option(None, help="Negate the --phrase-match filter (repeatable)."), + category: list[str] | None = typer.Option(None, help="Industry category filter (repeatable)."), + country: list[str] | None = typer.Option(None, help="ISO country code filter (repeatable)."), + state: list[str] | None = typer.Option(None, help="State or region filter (repeatable)."), + employee_range: str | None = typer.Option(None, help="Employee count range filter."), + revenue_range: str | None = typer.Option(None, help="Revenue range filter."), + variance: str | None = typer.Option( + None, help="Result diversity: LOW, MID_LOW, MEDIUM, MID_HIGH, HIGH, UNRESTRICTED." + ), + exclusion_query_id: list[str] | None = typer.Option(None, help=EXCLUSION_QUERY_ID_HELP), + param: list[str] | None = typer.Option(None, "--param", help=PARAM_HELP), + params_file: pathlib.Path | None = typer.Option(None, help=PARAMS_FILE_HELP), + max_companies: int = typer.Option(MAX_RECORDS, help="Stop once this many companies are in --out."), + page_size: int = typer.Option(MAX_RECORDS, min=MIN_RECORDS, max=MAX_RECORDS, help="Companies per discover call."), + run_name: str = typer.Option("bulk", help="Prefix for the exclusion lists created per page (-round-N)."), + out: pathlib.Path = typer.Option( + pathlib.Path("companies.csv"), help="CSV to append each page to; resumes if present." + ), + overwrite: bool = typer.Option(False, "--overwrite", help=OVERWRITE_HELP), + rate_limit: int = typer.Option(DEFAULT_RATE_LIMIT_PER_MINUTE, min=1, help=RATE_LIMIT_HELP), +) -> None: + """Walk the company index past 10,000 results, one exclusion list per page, appending to a CSV.""" + base = _drop_managed( + _merge_params( + param, + params_file, + icp_prompt=icp_prompt, + domain=domain, + phrase_match=phrase_match, + negate_phrase_match=negate_phrase_match, + category=category, + country=country, + state=state, + employee_range=employee_range, + revenue_range=revenue_range, + variance=variance, + ), + COMPANIES_MANAGED, + ) + build_request(DiscoverParams, {**base, "max_records": page_size}) # fail on a typo before any billable call + if not any(base.get(key) for key in ("icp_prompt", "icp_text", "domain")): + _log("warning: no icp_prompt / icp_text / seed --domain set - results are filter-only, unranked") + + seen: list[str] = [] + if out.exists() and not overwrite: + seen = read_domains_file(out) + _log(f"resuming: {len(seen):,} companies already in {out}") + seen_set = set(seen) + + client = _client(ctx) + limiter = _RateLimiter(rate_limit) + exclusion_ids: list[str] = list(exclusion_query_id or []) + inline_excludes: list[str] = [] + + def arm_exclusion(domains: list[str], label: str) -> None: + if len(domains) < MIN_EXCLUSION_LIST: + # Short tails ride on exclude_domain (100 max); once enough pile up they become a saved list too. + inline_excludes.extend(domains) + if len(inline_excludes) < MIN_EXCLUSION_LIST: + return + domains, label = list(inline_excludes), f"{label}-tails" + inline_excludes.clear() + request = CreateExclusionListRequest(query_name=label, domains=domains) + result = _call_with_retry(limiter, functools.partial(client.queries.create_exclusion_list, request)) + exclusion_ids.append(str(result.query_id)) + + if seen: + # --exclusion-query-id may be a suppression list unrelated to this CSV, so the file is always re-excluded. + supplied = len(exclusion_ids) + for index, batch in _chunked(seen, MAX_RECORDS): + arm_exclusion(batch, f"{run_name}-resume-{index}") + _log(f"re-armed {len(exclusion_ids) - supplied} exclusion list(s) from the existing file") + + added = 0 + rounds = 0 + with out.open("a" if seen else "w", newline="") as handle: + writer = csv.writer(handle) + if not seen: + writer.writerow(COMPANY_COLUMNS) + while len(seen) < max_companies: + rounds += 1 + want = min(page_size, max_companies - len(seen)) + request = build_request( + DiscoverParams, + { + **base, + "exclusion_query_id": exclusion_ids or None, + "exclude_domain": inline_excludes or None, + "max_records": max(MIN_RECORDS, want), + }, + ) + companies = _call_with_retry(limiter, functools.partial(client.discover, request)) + fresh = [c for c in companies if c.domain and c.domain not in seen_set][: max_companies - len(seen)] + for company in fresh: + assert company.domain is not None + seen.append(company.domain) + seen_set.add(company.domain) + writer.writerow( + [ + company.domain, + company.name, + company.address.country if company.address else None, + company.employees, + company.similarity, + ] + ) + handle.flush() + added += len(fresh) + _log(f"round {rounds}: +{len(fresh):,} net-new ({len(seen):,} total)") + if not fresh: + _log("no net-new companies left for this ICP - stopping") + break + if len(companies) < want: + _log("search returned less than requested - full set reached") + break + if len(seen) >= EXCLUSION_CAPACITY: + _log(f"warning: {len(seen):,} companies is at the exclusion capacity; slice the ICP into separate runs") + arm_exclusion([c.domain for c in fresh if c.domain], f"{run_name}-round-{rounds}") + + emit( + { + "companies": len(seen), + "new": added, + "rounds": rounds, + "exclusion_query_ids": exclusion_ids, + "out": str(out), + } + ) + + +# --------------------------------------------------------------------------- contacts / estimate + + +def _contact_filters( + param: list[str] | None, + params_file: pathlib.Path | None, + *, + icp_prompt: str | None, + summary: str | None, + negate_summary: str | None, + seniority: list[str] | None, + negate_seniority: list[str] | None, + department: list[str] | None, + negate_department: list[str] | None, + title: list[str] | None, + negate_title: list[str] | None, + person_country: list[str] | None, + person_state: list[str] | None, + has_email: bool, + exclusion_query_id: list[str] | None, +) -> dict[str, Any]: + return _drop_managed( + _merge_params( + param, + params_file, + icp_prompt=icp_prompt, + summary=summary, + negate_summary=negate_summary, + seniority=seniority, + negate_seniority=negate_seniority, + department=department, + negate_department=negate_department, + title=title, + negate_title=negate_title, + person_country=person_country, + person_state=person_state, + has_email=has_email, + exclusion_query_id=exclusion_query_id, + ), + CONTACTS_MANAGED, + ) + + +def _flatten(company: dict[str, Any], contact: dict[str, Any]) -> dict[str, Any]: + phones = contact.get("phone") or [] + first_phone = phones[0] if phones else None + linkedin = next((url for url in (contact.get("social_urls") or []) if "linkedin.com" in url), None) + name = contact.get("name") or "" + first, _, last = name.partition(" ") + return { + "persona_id": contact.get("persona_id"), + "first_name": contact.get("first_name") or first, + "last_name": contact.get("last_name") or last, + "name": name, + "title": contact.get("title"), + "seniority": contact.get("seniority"), + "department": contact.get("department"), + "email": contact.get("email"), + "email_validated": contact.get("email_validated"), + "phone": first_phone.get("phone") if isinstance(first_phone, dict) else first_phone, + "linkedin": linkedin, + "connections": contact.get("connections"), + "domain": company.get("domain"), + "company_name": company.get("name") or contact.get("company_name"), + "country": contact.get("country"), + "state": contact.get("state"), + "industry": ";".join(contact.get("industry") or []), + "employees": contact.get("employees") or company.get("employees"), + "revenue_range": contact.get("revenue_range") or company.get("revenue_range"), + "jobstart_date": contact.get("jobstart_date"), + } + + +def _checkpoint_fingerprint(domains: list[str], per_company: int, filters: dict[str, Any]) -> str: + """Identity of a contacts pull: same domains in the same order, same slice size, same filters.""" + payload = json.dumps( + {"domains": domains, "per_company": per_company, "filters": filters}, sort_keys=True, default=str + ) + return hashlib.sha256(payload.encode()).hexdigest()[:16] + + +def _read_checkpoint(path: pathlib.Path, fingerprint: str) -> set[int]: + """Slice indexes already pulled, refusing a checkpoint written for different inputs.""" + lines = path.read_text().split() + if not lines or lines[0] != f"fingerprint={fingerprint}": + raise typer.BadParameter( + f"{path} was written for a different domains file, --per-company or filters; " + "pass --overwrite to start clean or point --out elsewhere" + ) + return {int(line) for line in lines[1:]} + + +def _persona_ids_in(path: pathlib.Path) -> set[str]: + """Rows already in the CSV, so a rerun after a lost checkpoint never writes a persona twice.""" + if not path.exists(): + return set() + with path.open(newline="") as handle: + return {row["persona_id"] for row in csv.DictReader(handle) if row.get("persona_id")} + + +ICP_PROMPT_HELP = "Natural-language ICP prompt used to derive contact filters." +SUMMARY_HELP = "Filter by profile summary text (semantic search); ranks who comes back per company." +NEGATE_SUMMARY_HELP = "Exclude contacts matching this summary description." +HAS_EMAIL_HELP = "Only contacts with an email address (on by default)." + + +@app.command("estimate") +@handle_errors +def estimate_command( + ctx: typer.Context, + domains_file: pathlib.Path = typer.Option(..., "--domains-file", help=DOMAINS_FILE_HELP), + per_company: int = typer.Option(10, min=1, max=MAX_PER_COMPANY, help=PER_COMPANY_HELP), + icp_prompt: str | None = typer.Option(None, help=ICP_PROMPT_HELP), + summary: str | None = typer.Option(None, help=SUMMARY_HELP), + negate_summary: str | None = typer.Option(None, help=NEGATE_SUMMARY_HELP), + seniority: list[str] | None = typer.Option(None, help="Filter by seniority level (repeatable)."), + negate_seniority: list[str] | None = typer.Option(None, help="Exclude seniority levels (repeatable)."), + department: list[str] | None = typer.Option(None, help="Filter by department (repeatable)."), + negate_department: list[str] | None = typer.Option(None, help="Exclude departments (repeatable)."), + title: list[str] | None = typer.Option(None, help="Filter by job title (repeatable)."), + negate_title: list[str] | None = typer.Option(None, help="Exclude job titles (repeatable)."), + person_country: list[str] | None = typer.Option(None, help="Filter by contact country (repeatable)."), + person_state: list[str] | None = typer.Option(None, help="Filter by contact state/region (repeatable)."), + has_email: bool = typer.Option(True, "--has-email/--no-has-email", help=HAS_EMAIL_HELP), + exclusion_query_id: list[str] | None = typer.Option(None, help=EXCLUSION_QUERY_ID_HELP), + param: list[str] | None = typer.Option(None, "--param", help=PARAM_HELP), + params_file: pathlib.Path | None = typer.Option(None, help=PARAMS_FILE_HELP), + rate_limit: int = typer.Option( + COUNT_RATE_LIMIT_PER_MINUTE, min=1, help="Calls per minute on the free /contacts/count." + ), +) -> None: + """Size a contact pull for free: matching contacts across the domain list, and the cap at --per-company.""" + domains = read_domains_file(domains_file) + base = _contact_filters( + param, + params_file, + icp_prompt=icp_prompt, + summary=summary, + negate_summary=negate_summary, + seniority=seniority, + negate_seniority=negate_seniority, + department=department, + negate_department=negate_department, + title=title, + negate_title=negate_title, + person_country=person_country, + person_state=person_state, + has_email=has_email, + exclusion_query_id=exclusion_query_id, + ) + build_request(ContactsCountParams, {**base, "domain": domains[:1]}) + + client = _client(ctx) + limiter = _RateLimiter(rate_limit) + total = 0 + for index, batch in _chunked(domains, COUNT_SLICE): + request = build_request(ContactsCountParams, {**base, "domain": batch}) + count = _call_with_retry(limiter, functools.partial(client.contacts.count, request)) + total += count.count or 0 + _log(f" probed {min((index + 1) * COUNT_SLICE, len(domains)):,}/{len(domains):,} companies") + capped = min(total, len(domains) * per_company) + _log( + f"{len(domains):,} companies | {total:,} matching contacts | at most {capped:,} at {per_company}/company " + "(upper bound: the cap is applied to the total, a company with fewer matches pulls fewer)" + ) + emit( + {"companies": len(domains), "contacts_available": total, "contacts_capped": capped, "per_company": per_company} + ) + + +@app.command("contacts") +@handle_errors +def contacts_command( + ctx: typer.Context, + domains_file: pathlib.Path = typer.Option(..., "--domains-file", help=DOMAINS_FILE_HELP), + per_company: int = typer.Option(10, min=1, max=MAX_PER_COMPANY, help=PER_COMPANY_HELP), + icp_prompt: str | None = typer.Option(None, help=ICP_PROMPT_HELP), + summary: str | None = typer.Option(None, help=SUMMARY_HELP), + negate_summary: str | None = typer.Option(None, help=NEGATE_SUMMARY_HELP), + seniority: list[str] | None = typer.Option(None, help="Filter by seniority level (repeatable)."), + negate_seniority: list[str] | None = typer.Option(None, help="Exclude seniority levels (repeatable)."), + department: list[str] | None = typer.Option(None, help="Filter by department (repeatable)."), + negate_department: list[str] | None = typer.Option(None, help="Exclude departments (repeatable)."), + title: list[str] | None = typer.Option(None, help="Filter by job title (repeatable)."), + negate_title: list[str] | None = typer.Option(None, help="Exclude job titles (repeatable)."), + person_country: list[str] | None = typer.Option(None, help="Filter by contact country (repeatable)."), + person_state: list[str] | None = typer.Option(None, help="Filter by contact state/region (repeatable)."), + has_email: bool = typer.Option(True, "--has-email/--no-has-email", help=HAS_EMAIL_HELP), + exclusion_query_id: list[str] | None = typer.Option(None, help=EXCLUSION_QUERY_ID_HELP), + param: list[str] | None = typer.Option(None, "--param", help=PARAM_HELP), + params_file: pathlib.Path | None = typer.Option(None, help=PARAMS_FILE_HELP), + out: pathlib.Path = typer.Option( + pathlib.Path("contacts.csv"), help="CSV to append to; .checkpoint tracks slices." + ), + overwrite: bool = typer.Option(False, "--overwrite", help=OVERWRITE_HELP), + rate_limit: int = typer.Option(DEFAULT_RATE_LIMIT_PER_MINUTE, min=1, help=RATE_LIMIT_HELP), +) -> None: + """Pull up to --per-company contacts for every domain in the file into one CSV, resumable by slice.""" + domains = read_domains_file(domains_file) + per_call = max(1, MAX_RECORDS // per_company) + slices = list(_chunked(domains, per_call)) + base = _contact_filters( + param, + params_file, + icp_prompt=icp_prompt, + summary=summary, + negate_summary=negate_summary, + seniority=seniority, + negate_seniority=negate_seniority, + department=department, + negate_department=negate_department, + title=title, + negate_title=negate_title, + person_country=person_country, + person_state=person_state, + has_email=has_email, + exclusion_query_id=exclusion_query_id, + ) + build_request(ContactFilters, {**base, "domain": domains[:1], "results_by_company": per_company}) + + checkpoint = out.with_suffix(out.suffix + ".checkpoint") + if overwrite: + checkpoint.unlink(missing_ok=True) + out.unlink(missing_ok=True) + fingerprint = _checkpoint_fingerprint(domains, per_company, base) + done: set[int] = set() + if checkpoint.exists(): + done = _read_checkpoint(checkpoint, fingerprint) + _log(f"resuming: {len(done)}/{len(slices)} slices already pulled") + todo = [(index, batch) for index, batch in slices if index not in done] + _log(f"{len(domains):,} companies | {len(slices)} slice(s) of {per_call} | {len(todo)} to run") + + client = _client(ctx) + limiter = _RateLimiter(rate_limit) + seen_personas = _persona_ids_in(out) + written = 0 + with out.open("a", newline="") as handle, checkpoint.open("a") as checkpoint_handle: + writer = csv.DictWriter(handle, fieldnames=CONTACT_COLUMNS, extrasaction="ignore") + if handle.tell() == 0: + writer.writeheader() + if checkpoint_handle.tell() == 0: + checkpoint_handle.write(f"fingerprint={fingerprint}\n") + for index, batch in todo: + request = build_request( + ContactFilters, + { + **base, + "domain": batch, + "results_by_company": per_company, + "max_records": max(MIN_RECORDS, min(MAX_RECORDS, len(batch) * per_company)), + }, + ) + payload = _call_with_retry(limiter, functools.partial(client.contacts.discover, request)).to_dict() + rows = [ + _flatten(company, contact) + for company in (payload.get("results") or {}).values() + for contact in (company.get("contacts") or []) + ] + # A contact without a persona_id cannot be deduplicated by id, so it is always written. + fresh = [row for row in rows if row["persona_id"] is None or str(row["persona_id"]) not in seen_personas] + seen_personas.update(str(row["persona_id"]) for row in fresh if row["persona_id"] is not None) + writer.writerows(fresh) + handle.flush() + checkpoint_handle.write(f"{index}\n") + checkpoint_handle.flush() + written += len(fresh) + _log(f" slice {index + 1}/{len(slices)}: {len(fresh):,} contacts ({written:,} this run)") + + emit( + { + "companies": len(domains), + "slices": len(slices), + "slices_run": len(todo), + "contacts": written, + "out": str(out), + "checkpoint": str(checkpoint), + } + ) diff --git a/packages/discolike-cli/src/discolike_cli/contacts.py b/packages/discolike-cli/src/discolike_cli/contacts.py index 91d1c9c..be81c1e 100644 --- a/packages/discolike-cli/src/discolike_cli/contacts.py +++ b/packages/discolike-cli/src/discolike_cli/contacts.py @@ -12,10 +12,15 @@ from discolike.requests import ContactsLookupParams from discolike.requests import ContactsMatchParams from discolike.requests import ContactsSearchParams +from discolike_cli._help import ContractCommand +from discolike_cli._help import epilog +from discolike_cli._inputs import merge_domains from discolike_cli._output import build_request from discolike_cli._output import emit from discolike_cli._output import handle_errors from discolike_cli._output import run_job +from discolike_cli.discover import DOMAINS_FILE_HELP +from discolike_cli.discover import PARAMS_FILE_HELP from discolike_cli.discover import _merge_params DEFAULT_WAIT_TIMEOUT_SECONDS = 900.0 @@ -50,7 +55,7 @@ ) -@app.command("search") +@app.command("search", cls=ContractCommand, epilog=epilog("contacts search")) @handle_errors def search_command( ctx: typer.Context, @@ -66,6 +71,7 @@ def search_command( negate_summary: str | None = typer.Option(None, help=NEGATE_SUMMARY_HELP), skills: list[str] | None = typer.Option(None, help=SKILLS_HELP), domain: list[str] | None = typer.Option(None, help="Filter by company domain (repeatable)."), + domains_file: pathlib.Path | None = typer.Option(None, help=DOMAINS_FILE_HELP), person_country: list[str] | None = typer.Option(None, help="Filter by contact country (repeatable)."), negate_person_country: list[str] | None = typer.Option(None, help="Exclude contact countries (repeatable)."), person_state: list[str] | None = typer.Option(None, help=PERSON_STATE_HELP), @@ -100,12 +106,14 @@ def search_command( consensus: int | None = typer.Option(None, help=CONSENSUS_HELP), fmt: str | None = typer.Option(None, "--format", help=FORMAT_HELP), param: list[str] | None = typer.Option(None, "--param", help=PARAM_HELP), + params_file: pathlib.Path | None = typer.Option(None, help=PARAMS_FILE_HELP), ) -> None: """Search contacts matching the given filters.""" from discolike_cli.main import get_client kwargs = _merge_params( param, + params_file, icp_prompt=icp_prompt, seniority=seniority, negate_seniority=negate_seniority, @@ -117,7 +125,7 @@ def search_command( summary=summary, negate_summary=negate_summary, skills=skills, - domain=domain, + domain=merge_domains(domain, domains_file), person_country=person_country, negate_person_country=negate_person_country, person_state=person_state, @@ -164,6 +172,7 @@ def count_command( negate_summary: str | None = typer.Option(None, help=NEGATE_SUMMARY_HELP), skills: list[str] | None = typer.Option(None, help=SKILLS_HELP), domain: list[str] | None = typer.Option(None, help="Filter by company domain (repeatable)."), + domains_file: pathlib.Path | None = typer.Option(None, help=DOMAINS_FILE_HELP), person_country: list[str] | None = typer.Option(None, help="Filter by contact country (repeatable)."), negate_person_country: list[str] | None = typer.Option(None, help="Exclude contact countries (repeatable)."), person_state: list[str] | None = typer.Option(None, help=PERSON_STATE_HELP), @@ -198,12 +207,14 @@ def count_command( consensus: int | None = typer.Option(None, help=CONSENSUS_HELP), fmt: str | None = typer.Option(None, "--format", help=FORMAT_HELP), param: list[str] | None = typer.Option(None, "--param", help=PARAM_HELP), + params_file: pathlib.Path | None = typer.Option(None, help=PARAMS_FILE_HELP), ) -> None: """Count contacts matching the given filters.""" from discolike_cli.main import get_client kwargs = _merge_params( param, + params_file, icp_prompt=icp_prompt, seniority=seniority, negate_seniority=negate_seniority, @@ -215,7 +226,7 @@ def count_command( summary=summary, negate_summary=negate_summary, skills=skills, - domain=domain, + domain=merge_domains(domain, domains_file), person_country=person_country, negate_person_country=negate_person_country, person_state=person_state, @@ -332,6 +343,7 @@ def discover_command( negate_summary: str | None = typer.Option(None, help=NEGATE_SUMMARY_HELP), skills: list[str] | None = typer.Option(None, help=SKILLS_HELP), domain: list[str] | None = typer.Option(None, help="Filter by company domain (repeatable)."), + domains_file: pathlib.Path | None = typer.Option(None, help=DOMAINS_FILE_HELP), person_country: list[str] | None = typer.Option(None, help="Filter by contact country (repeatable)."), negate_person_country: list[str] | None = typer.Option(None, help="Exclude contact countries (repeatable)."), person_state: list[str] | None = typer.Option(None, help=PERSON_STATE_HELP), @@ -370,12 +382,14 @@ def discover_command( consensus: int | None = typer.Option(None, "--consensus", help="Consensus threshold for discovered contacts."), fmt: str | None = typer.Option(None, "--format", help=FORMAT_HELP), param: list[str] | None = typer.Option(None, "--param", help=PARAM_HELP), + params_file: pathlib.Path | None = typer.Option(None, help=PARAMS_FILE_HELP), ) -> None: """Discover contacts grouped by company for the given filters.""" from discolike_cli.main import get_client kwargs = _merge_params( param, + params_file, icp_prompt=icp_prompt, seniority=seniority, negate_seniority=negate_seniority, @@ -387,7 +401,7 @@ def discover_command( summary=summary, negate_summary=negate_summary, skills=skills, - domain=domain, + domain=merge_domains(domain, domains_file), person_country=person_country, negate_person_country=negate_person_country, person_state=person_state, @@ -418,12 +432,13 @@ def discover_command( emit(get_client(ctx).contacts.discover(build_request(ContactFilters, kwargs)), fmt=fmt) -@app.command("generate") +@app.command("generate", cls=ContractCommand, epilog=epilog("contacts generate")) @handle_errors def generate_command( ctx: typer.Context, icp_text: str = typer.Option(..., "--icp-text", help="ICP description used to generate contacts."), - domain: list[str] = typer.Option(..., "--domain", help="Target company domain (repeatable)."), + domain: list[str] | None = typer.Option(None, "--domain", help="Target company domain (repeatable)."), + domains_file: pathlib.Path | None = typer.Option(None, "--domains-file", help=DOMAINS_FILE_HELP), full_domain: list[str] | None = typer.Option( None, "--full-domain", help="Domain to send as full_domains to the generation job (repeatable)." ), @@ -444,6 +459,11 @@ def generate_command( max_company_records: int | None = typer.Option( None, "--max-company-records", help="Maximum company records to process." ), + find_emails: bool = typer.Option( + False, + "--find-emails", + help="Run the email finder over named, email-less rows before the job completes (found addresses bill).", + ), wait: bool = typer.Option(False, "--wait", help=WAIT_HELP), timeout: float = typer.Option(DEFAULT_WAIT_TIMEOUT_SECONDS, "--timeout", help=TIMEOUT_HELP), fmt: str | None = typer.Option(None, "--format", help=FORMAT_HELP), @@ -451,12 +471,15 @@ def generate_command( """Generate contacts for target domains from an ICP description (async job).""" from discolike_cli.main import get_client + if not domain and domains_file is None: + raise typer.BadParameter("Provide --domain or --domains-file") + request = build_request( ContactGenerateRequest, _merge_params( None, icp_text=icp_text, - domains=domain, + domains=merge_domains(domain, domains_file), full_domains=full_domain, partial_domains=partial_domain, context_mode=context_mode, @@ -465,6 +488,7 @@ def generate_command( search_context_size=search_context_size, max_contacts_per_domain=max_contacts_per_domain, max_company_records=max_company_records, + find_emails=find_emails or None, ), ) run_job(get_client(ctx).contacts.generate(request), wait=wait, timeout=timeout, fmt=fmt) diff --git a/packages/discolike-cli/src/discolike_cli/discogen.py b/packages/discolike-cli/src/discolike_cli/discogen.py index fbb33be..b1ff0f6 100644 --- a/packages/discolike-cli/src/discolike_cli/discogen.py +++ b/packages/discolike-cli/src/discolike_cli/discogen.py @@ -1,16 +1,21 @@ from __future__ import annotations import enum +import pathlib import typer from discolike._jobs import Job from discolike.requests import DiscoGenPersonaProcessRequest from discolike.requests import DiscoGenProcessRequest +from discolike_cli._help import ContractCommand +from discolike_cli._help import epilog +from discolike_cli._inputs import merge_domains from discolike_cli._output import build_request from discolike_cli._output import emit from discolike_cli._output import handle_errors from discolike_cli._output import run_job +from discolike_cli.discover import DOMAINS_FILE_HELP from discolike_cli.discover import _merge_params DEFAULT_WAIT_TIMEOUT_SECONDS = 900.0 @@ -23,10 +28,15 @@ "segment jobs 'segment', contact bulk-match 'contactmatch')." ) QUERY_HELP = "Research query to run." -INTEGRATION_ID_HELP = "Integration ID to use for the run." +INTEGRATION_ID_HELP = ( + "Integration ID to use for the run, or 'native-icp' to score an ICP validation prompt with " + "DiscoLike's own model at no LLM cost." +) WEB_SEARCH_HELP = "Toggle web search during research." CONTEXT_MODE_HELP = "Context mode; see docs.discolike.com." INCLUDE_X_SEARCH_HELP = "Toggle including X search in the research." +TYPED_COLUMNS_HELP = "Let the detector answer yes/no, fixed-set and scale columns with a TypeSafe judgment model." +INCLUDE_CONFIDENCE_HELP = "Add a confidence column beside each typed column." SEARCH_PROVIDER_ID_HELP = "Search provider ID to use for web search." SEARCH_CONTEXT_SIZE_HELP = "Search context size; see docs.discolike.com." TASK_ID_HELP = "Task ID returned when the job was started." @@ -41,18 +51,23 @@ class TaskFamily(str, enum.Enum): segment = "segment" -@app.command("run") +@app.command("run", cls=ContractCommand, epilog=epilog("discogen run")) @handle_errors def run_command( ctx: typer.Context, query: str = typer.Option(..., "--query", help=QUERY_HELP), - domain: list[str] = typer.Option(..., "--domain", help="Company domain to research (repeatable)."), + domain: list[str] | None = typer.Option(None, "--domain", help="Company domain to research (repeatable)."), + domains_file: pathlib.Path | None = typer.Option(None, "--domains-file", help=DOMAINS_FILE_HELP), integration_id: str | None = typer.Option(None, "--integration-id", help=INTEGRATION_ID_HELP), web_search: bool | None = typer.Option(None, "--web-search/--no-web-search", help=WEB_SEARCH_HELP), context_mode: str | None = typer.Option(None, "--context-mode", help=CONTEXT_MODE_HELP), include_x_search: bool | None = typer.Option( None, "--include-x-search/--no-include-x-search", help=INCLUDE_X_SEARCH_HELP ), + typed_columns: bool | None = typer.Option(None, "--typed-columns/--no-typed-columns", help=TYPED_COLUMNS_HELP), + include_confidence: bool | None = typer.Option( + None, "--include-confidence/--no-include-confidence", help=INCLUDE_CONFIDENCE_HELP + ), search_provider_id: str | None = typer.Option(None, "--search-provider-id", help=SEARCH_PROVIDER_ID_HELP), search_context_size: str | None = typer.Option(None, "--search-context-size", help=SEARCH_CONTEXT_SIZE_HELP), wait: bool = typer.Option(False, "--wait", help=WAIT_HELP), @@ -62,16 +77,21 @@ def run_command( """Run a DiscoGen research query across company domains (async job).""" from discolike_cli.main import get_client + if not domain and domains_file is None: + raise typer.BadParameter("Provide --domain or --domains-file") + request = build_request( DiscoGenProcessRequest, _merge_params( None, query=query, - domains=domain, + domains=merge_domains(domain, domains_file), integration_id=integration_id, web_search=web_search, context_mode=context_mode, include_x_search=include_x_search, + typed_columns=typed_columns, + include_confidence=include_confidence, search_provider_id=search_provider_id, search_context_size=search_context_size, ), @@ -91,6 +111,10 @@ def run_personas_command( include_x_search: bool | None = typer.Option( None, "--include-x-search/--no-include-x-search", help=INCLUDE_X_SEARCH_HELP ), + typed_columns: bool | None = typer.Option(None, "--typed-columns/--no-typed-columns", help=TYPED_COLUMNS_HELP), + include_confidence: bool | None = typer.Option( + None, "--include-confidence/--no-include-confidence", help=INCLUDE_CONFIDENCE_HELP + ), search_provider_id: str | None = typer.Option(None, "--search-provider-id", help=SEARCH_PROVIDER_ID_HELP), search_context_size: str | None = typer.Option(None, "--search-context-size", help=SEARCH_CONTEXT_SIZE_HELP), wait: bool = typer.Option(False, "--wait", help=WAIT_HELP), @@ -110,6 +134,8 @@ def run_personas_command( web_search=web_search, context_mode=context_mode, include_x_search=include_x_search, + typed_columns=typed_columns, + include_confidence=include_confidence, search_provider_id=search_provider_id, search_context_size=search_context_size, ), @@ -138,7 +164,7 @@ def _build_job(ctx: typer.Context, family: TaskFamily, task_id: str) -> Job: return Job(client._transport, task_family=family.value, task_id=task_id) -@app.command("status") +@app.command("status", cls=ContractCommand, epilog=epilog("discogen status")) @handle_errors def status_command( ctx: typer.Context, diff --git a/packages/discolike-cli/src/discolike_cli/discover.py b/packages/discolike-cli/src/discolike_cli/discover.py index b9c2df1..7d743f8 100644 --- a/packages/discolike-cli/src/discolike_cli/discover.py +++ b/packages/discolike-cli/src/discolike_cli/discover.py @@ -1,11 +1,14 @@ from __future__ import annotations +import pathlib from typing import Any import typer from discolike.requests import CountParams from discolike.requests import DiscoverParams +from discolike_cli._inputs import merge_domains +from discolike_cli._inputs import read_params_file from discolike_cli._output import build_request from discolike_cli._output import emit from discolike_cli._output import handle_errors @@ -15,6 +18,18 @@ FORMAT_HELP = "Output format: json or table (table auto-selected on a TTY; falls back to JSON for non-tabular data)." PARAM_HELP = "Extra API parameter as KEY=VALUE (comma-separates into a list); see docs.discolike.com" +BBOX_HELP = "Bounding box as min_lat,min_lon,max_lat,max_lon (repeatable). Longitudes may wrap the antimeridian." +GEO_HELP = "Circular area as lat,lon or lat,lon,radius, e.g. 30.27,-97.74,30mi (repeatable). Radius defaults to 50km." +SHAPES_ARE_ORED_HELP = ( + "Every --geo circle, every --bbox and the --lat/--lon/--radius centre are OR'd together, up to 10 in total." +) +PARAMS_FILE_HELP = ( + "JSON object of API parameter names (e.g. an app form copied over); --param and flags override its values." +) +DOMAINS_FILE_HELP = "CSV with a 'domain' column, or one domain per line; merged with --domain." +EXCLUDE_DOMAINS_FILE_HELP = ( + "CSV with a 'domain' column, or one domain per line; merged with --exclude-domain (100 max)." +) SUBDOMAIN_HELP = "Limit results to this subdomain, up to 20, each at least 3 characters (repeatable)." NEGATE_SUBDOMAIN_HELP = "Exclude this subdomain, up to 20, each at least 3 characters (repeatable)." START_DATE_HELP = "Minimum company start date (YYYY-MM-DD) or range (YYYY-MM-DD,YYYY-MM-DD)." @@ -39,8 +54,14 @@ def _parse_param(raw: str) -> tuple[str, str | list[str]]: return key, value -def _merge_params(param: list[str] | None, **options: Any) -> dict[str, Any]: # noqa: ANN401 -- forwarded as a dict to build_request - kwargs: dict[str, Any] = dict(_parse_param(raw) for raw in param or []) +def _merge_params( + param: list[str] | None, + params_file: pathlib.Path | None = None, + **options: Any, # noqa: ANN401 -- forwarded as a dict to build_request +) -> dict[str, Any]: + """Precedence, lowest to highest: --params-file, --param KEY=VALUE, first-class flags.""" + kwargs: dict[str, Any] = read_params_file(params_file) if params_file is not None else {} + kwargs.update(_parse_param(raw) for raw in param or []) kwargs.update({key: value for key, value in options.items() if value is not None}) return kwargs @@ -54,6 +75,13 @@ def discover_command( negate_phrase_match: list[str] | None = typer.Option(None, help="Negate the --phrase-match filter (repeatable)."), category: list[str] | None = typer.Option(None, help="Industry category filter (repeatable)."), negate_category: list[str] | None = typer.Option(None, help="Negate the --category filter (repeatable)."), + sub_industry: list[str] | None = typer.Option(None, help="Sub-industry filter, bare or PARENT/SUB (repeatable)."), + negate_sub_industry: list[str] | None = typer.Option(None, help="Negate the --sub-industry filter (repeatable)."), + lat: float | None = typer.Option(None, help="Latitude of the search centre; requires --lon."), + lon: float | None = typer.Option(None, help="Longitude of the search centre; requires --lat."), + radius: str | None = typer.Option(None, help="Radius around --lat/--lon, e.g. 50km or 30mi. Default 50km."), + geo: list[str] | None = typer.Option(None, help=f"{GEO_HELP} {SHAPES_ARE_ORED_HELP}"), + bbox: list[str] | None = typer.Option(None, help=f"{BBOX_HELP} {SHAPES_ARE_ORED_HELP}"), country: list[str] | None = typer.Option(None, help="ISO country code filter (repeatable)."), negate_country: list[str] | None = typer.Option(None, help="Negate the --country filter (repeatable)."), state: list[str] | None = typer.Option(None, help="State or region filter (repeatable)."), @@ -104,6 +132,7 @@ def discover_command( None, "--auto-phrase-match/--no-auto-phrase-match", help="Auto-generate phrase matches from ICP text." ), exclude_domain: list[str] | None = typer.Option(None, help="Domain to exclude from results (repeatable)."), + exclude_domains_file: pathlib.Path | None = typer.Option(None, help=EXCLUDE_DOMAINS_FILE_HELP), inclusion_query_id: list[str] | None = typer.Option( None, help="Saved query ID whose domains are included (repeatable); requires the STARTER plan." ), @@ -114,6 +143,7 @@ def discover_command( offset: int | None = typer.Option(None, help="Number of records to skip for pagination."), fmt: str | None = typer.Option(None, "--format", help=FORMAT_HELP), param: list[str] | None = typer.Option(None, "--param", help=PARAM_HELP), + params_file: pathlib.Path | None = typer.Option(None, help=PARAMS_FILE_HELP), ) -> None: """Discover companies matching your ICP and filters.""" from discolike_cli.main import get_client @@ -122,12 +152,20 @@ def discover_command( DiscoverParams, _merge_params( param, + params_file, icp_prompt=icp_prompt, domain=domain, phrase_match=phrase_match, negate_phrase_match=negate_phrase_match, category=category, negate_category=negate_category, + sub_industry=sub_industry, + negate_sub_industry=negate_sub_industry, + lat=lat, + lon=lon, + radius=radius, + geo=geo, + bbox=bbox, country=country, negate_country=negate_country, state=state, @@ -157,7 +195,7 @@ def discover_command( include_search_domains=include_search_domains, auto_icp_text=auto_icp_text, auto_phrase_match=auto_phrase_match, - exclude_domain=exclude_domain, + exclude_domain=merge_domains(exclude_domain, exclude_domains_file), inclusion_query_id=inclusion_query_id, exclusion_query_id=exclusion_query_id, max_records=max_records, @@ -174,6 +212,13 @@ def count_command( negate_phrase_match: list[str] | None = typer.Option(None, help="Negate the --phrase-match filter (repeatable)."), category: list[str] | None = typer.Option(None, help="Industry category filter (repeatable)."), negate_category: list[str] | None = typer.Option(None, help="Negate the --category filter (repeatable)."), + sub_industry: list[str] | None = typer.Option(None, help="Sub-industry filter, bare or PARENT/SUB (repeatable)."), + negate_sub_industry: list[str] | None = typer.Option(None, help="Negate the --sub-industry filter (repeatable)."), + lat: float | None = typer.Option(None, help="Latitude of the search centre; requires --lon."), + lon: float | None = typer.Option(None, help="Longitude of the search centre; requires --lat."), + radius: str | None = typer.Option(None, help="Radius around --lat/--lon, e.g. 50km or 30mi. Default 50km."), + geo: list[str] | None = typer.Option(None, help=f"{GEO_HELP} {SHAPES_ARE_ORED_HELP}"), + bbox: list[str] | None = typer.Option(None, help=f"{BBOX_HELP} {SHAPES_ARE_ORED_HELP}"), country: list[str] | None = typer.Option(None, help="ISO country code filter (repeatable)."), negate_country: list[str] | None = typer.Option(None, help="Negate the --country filter (repeatable)."), state: list[str] | None = typer.Option(None, help="State or region filter (repeatable)."), @@ -201,6 +246,7 @@ def count_command( ), fmt: str | None = typer.Option(None, "--format", help=FORMAT_HELP), param: list[str] | None = typer.Option(None, "--param", help=PARAM_HELP), + params_file: pathlib.Path | None = typer.Option(None, help=PARAMS_FILE_HELP), ) -> None: """Count companies matching the given filters.""" from discolike_cli.main import get_client @@ -209,10 +255,18 @@ def count_command( CountParams, _merge_params( param, + params_file, phrase_match=phrase_match, negate_phrase_match=negate_phrase_match, category=category, negate_category=negate_category, + sub_industry=sub_industry, + negate_sub_industry=negate_sub_industry, + lat=lat, + lon=lon, + radius=radius, + geo=geo, + bbox=bbox, country=country, negate_country=negate_country, state=state, diff --git a/packages/discolike-cli/src/discolike_cli/enrich.py b/packages/discolike-cli/src/discolike_cli/enrich.py index d6144b8..89cd238 100644 --- a/packages/discolike-cli/src/discolike_cli/enrich.py +++ b/packages/discolike-cli/src/discolike_cli/enrich.py @@ -8,6 +8,7 @@ from discolike.requests import SegmentFileParams from discolike.requests import SegmentParams from discolike.requests import ValidateIcpRequest +from discolike_cli._inputs import read_domains_file from discolike_cli._output import build_request from discolike_cli._output import emit from discolike_cli._output import handle_errors @@ -20,6 +21,10 @@ WAIT_HELP = "Block until the job finishes, streaming progress to stderr." TIMEOUT_HELP = "Max seconds to wait with --wait." QUERY_ID_HELP = "Saved query ID whose domains are included alongside the file/--domain ones (repeatable)." +INTEGRATION_ID_HELP = ( + "Integration ID to use for the validation, or 'native-icp' to score with DiscoLike's own " + "ICP-fit model at no LLM cost (no LLM key and no web search on that run)." +) @handle_errors @@ -28,12 +33,13 @@ def validate_icp_command( icp: str = typer.Option(..., "--icp", help="ICP definition text to validate the domains against."), domain: list[str] | None = typer.Option(None, "--domain", help="Domain to validate (repeatable)."), file: pathlib.Path | None = typer.Option( - None, "--file", help="Text file with one domain per line (instead of --domain)." + None, + "--domains-file", + "--file", + help="CSV with a 'domain' column, or one domain per line (instead of --domain).", ), context_mode: str | None = typer.Option(None, "--context-mode", help="Context mode; see docs.discolike.com."), - integration_id: str | None = typer.Option( - None, "--integration-id", help="Integration ID to use for the validation." - ), + integration_id: str | None = typer.Option(None, "--integration-id", help=INTEGRATION_ID_HELP), web_search: bool | None = typer.Option( None, "--web-search/--no-web-search", help="Toggle web search during validation." ), @@ -48,13 +54,9 @@ def validate_icp_command( from discolike_cli.main import get_client if (not domain) == (file is None): - raise typer.BadParameter("Provide exactly one of --domain or --file") + raise typer.BadParameter("Provide exactly one of --domain or --domains-file") - if file is not None: - domains = [line.strip() for line in file.read_text().splitlines() if line.strip()] - else: - assert domain is not None - domains = domain + domains = read_domains_file(file) if file is not None else domain request = build_request( ValidateIcpRequest, diff --git a/packages/discolike-cli/src/discolike_cli/main.py b/packages/discolike-cli/src/discolike_cli/main.py index d655aca..42457bd 100644 --- a/packages/discolike-cli/src/discolike_cli/main.py +++ b/packages/discolike-cli/src/discolike_cli/main.py @@ -1,14 +1,21 @@ from __future__ import annotations +import sys from importlib.metadata import version as package_version from typing import Any import typer +from typer._click.exceptions import Abort +from typer._click.exceptions import ClickException +from typer._click.exceptions import NoArgsIsHelpError +from typer._click.exceptions import UsageError from discolike import Discolike +from discolike import ValidationError from discolike import __version__ as sdk_version from discolike_cli import account from discolike_cli import auth +from discolike_cli import bulk from discolike_cli import company from discolike_cli import contacts from discolike_cli import discogen @@ -19,6 +26,11 @@ from discolike_cli import providers from discolike_cli import queries from discolike_cli import signup +from discolike_cli._help import MAIN_EPILOG +from discolike_cli._help import ContractCommand +from discolike_cli._help import ContractGroup +from discolike_cli._help import epilog +from discolike_cli._output import fail app = typer.Typer( name="discolike", @@ -30,8 +42,15 @@ "company names to domains, and find the right contacts — from your terminal.\n\n" "Docs: https://docs.discolike.com · Keys: https://app.discolike.com/account/management/keys" ), + epilog=MAIN_EPILOG, + cls=ContractGroup, ) + +def _stdout_is_tty() -> bool: + return sys.stdout.isatty() + + build_client = Discolike @@ -45,7 +64,11 @@ def main( version: bool = typer.Option(False, "--version", help="Print CLI and SDK versions and exit."), ) -> None: if version: - typer.echo(f"🪩 discolike-cli {package_version('discolike-cli')} (discolike {sdk_version})") + cli_version = package_version("discolike-cli") + if _stdout_is_tty(): + typer.echo(f"🪩 discolike-cli {cli_version} (discolike {sdk_version})") + else: + typer.echo(cli_version) raise typer.Exit ctx.obj = {"api_key": api_key, "base_url": base_url} @@ -56,6 +79,7 @@ def get_client(ctx: typer.Context) -> Discolike: app.add_typer(auth.app, name="auth") +app.add_typer(bulk.app, name="bulk") app.add_typer(company.app, name="company") app.add_typer(contacts.app, name="contacts") app.add_typer(discogen.app, name="discogen") @@ -64,11 +88,40 @@ def get_client(ctx: typer.Context) -> Discolike: app.add_typer(account.app, name="account") app.add_typer(providers.search_providers_app, name="search-providers") app.add_typer(providers.llm_providers_app, name="llm-providers") -app.command(name="discover")(discover.discover_command) -app.command(name="count")(discover.count_command) -app.command(name="match")(match.match_command) -app.command(name="extract")(company.extract_command) -app.command(name="validate-icp")(enrich.validate_icp_command) -app.command(name="append")(enrich.append_command) -app.command(name="segment")(enrich.segment_command) -app.command(name="signup")(signup.signup_command) +KEYBOARD_INTERRUPT_EXIT_CODE = 130 + + +def run(argv: list[str] | None = None) -> None: + """Console entry point: every failure, parser errors included, honours the JSON envelope. + + Click's standalone mode prints its own usage text on a bad flag or value; running the app + with ``standalone_mode=False`` lets us route those through ``fail()`` so agents always see + ``{"code": "validation_error", ...}`` on stderr with exit 2. Help, ``--version``, and the + bare-invocation help screen keep click's behaviour. + """ + try: + result = app(args=argv, prog_name="discolike", standalone_mode=False) + except NoArgsIsHelpError as exc: + exc.show() + sys.exit(exc.exit_code) + except UsageError as exc: + message = exc.format_message() + if exc.ctx is not None: + message = f"{message} (see `{exc.ctx.command_path} --help`)" + sys.exit(fail(ValidationError(message)).exit_code) + except ClickException as exc: + exc.show() + sys.exit(exc.exit_code) + except (Abort, KeyboardInterrupt): + sys.exit(KEYBOARD_INTERRUPT_EXIT_CODE) + sys.exit(result if isinstance(result, int) else 0) + + +app.command(name="discover", cls=ContractCommand, epilog=epilog("discover"))(discover.discover_command) +app.command(name="count", cls=ContractCommand, epilog=epilog("count"))(discover.count_command) +app.command(name="match", cls=ContractCommand, epilog=epilog("match"))(match.match_command) +app.command(name="extract", cls=ContractCommand, epilog=epilog("extract"))(company.extract_command) +app.command(name="validate-icp", cls=ContractCommand, epilog=epilog("validate-icp"))(enrich.validate_icp_command) +app.command(name="append", cls=ContractCommand, epilog=epilog("append"))(enrich.append_command) +app.command(name="segment", cls=ContractCommand, epilog=epilog("segment"))(enrich.segment_command) +app.command(name="signup", cls=ContractCommand, epilog=epilog("signup"))(signup.signup_command) diff --git a/packages/discolike-cli/src/discolike_cli/queries.py b/packages/discolike-cli/src/discolike_cli/queries.py index 71a23f8..bd35625 100644 --- a/packages/discolike-cli/src/discolike_cli/queries.py +++ b/packages/discolike-cli/src/discolike_cli/queries.py @@ -11,9 +11,13 @@ from discolike.requests import QueriesListParams from discolike.requests import SaveResultsRequest from discolike.requests import UpdateQueryRequest +from discolike_cli._help import ContractCommand +from discolike_cli._help import epilog +from discolike_cli._inputs import merge_domains from discolike_cli._output import build_request from discolike_cli._output import emit from discolike_cli._output import handle_errors +from discolike_cli.discover import DOMAINS_FILE_HELP from discolike_cli.discover import _merge_params FORMAT_HELP = "Output format: json or table (table auto-selected on a TTY; falls back to JSON for non-tabular data)." @@ -48,12 +52,13 @@ def list_command( emit(get_client(ctx).queries.list(request), fmt=fmt) -@app.command("create-exclusion-list") +@app.command("create-exclusion-list", cls=ContractCommand, epilog=epilog("queries create-exclusion-list")) @handle_errors def create_exclusion_list_command( ctx: typer.Context, name: str = typer.Option(..., "--name", help="Name for the new exclusion list."), domain: list[str] | None = typer.Option(None, "--domain", help="Domain to exclude (repeatable)."), + domains_file: Path | None = typer.Option(None, "--domains-file", help=DOMAINS_FILE_HELP), persona_id: list[int] | None = typer.Option(None, "--persona-id", help="Persona ID to exclude (repeatable)."), tag: list[str] | None = typer.Option(None, "--tag", help="Tag to attach to the list (repeatable)."), ) -> None: @@ -62,7 +67,9 @@ def create_exclusion_list_command( request = build_request( CreateExclusionListRequest, - _merge_params(None, query_name=name, domains=domain, persona_ids=persona_id, tags=tag), + _merge_params( + None, query_name=name, domains=merge_domains(domain, domains_file), persona_ids=persona_id, tags=tag + ), ) emit(get_client(ctx).queries.create_exclusion_list(request)) @@ -92,6 +99,8 @@ def save_results_command( data = json.loads(input_path.read_text()) except FileNotFoundError as exc: raise typer.BadParameter(f"--input file not found: {input_path}") from exc + except (OSError, UnicodeDecodeError) as exc: + raise typer.BadParameter(f"--input file {input_path} could not be read: {exc}") from exc except json.JSONDecodeError as exc: raise typer.BadParameter(f"--input file {input_path} must contain valid JSON: {exc}") from exc diff --git a/packages/discolike-cli/tests/test_auth.py b/packages/discolike-cli/tests/test_auth.py index f7b3de9..4a6819d 100644 --- a/packages/discolike-cli/tests/test_auth.py +++ b/packages/discolike-cli/tests/test_auth.py @@ -33,8 +33,9 @@ def test_login_with_api_key_option_verifies_and_saves(install_build_client: Call assert json.loads(config_path().read_text())["api_key"] == "dk-1" mode = stat.S_IMODE(config_path().stat().st_mode) assert mode == 0o600 - payload = json.loads(result.stderr) + payload = json.loads(result.stdout) assert payload["logged_in"] is True + assert result.stderr == "" def test_login_prompts_for_key_when_not_given(install_build_client: Callable[[Handler], None]) -> None: @@ -118,8 +119,8 @@ def test_cli_version_flag() -> None: result = runner.invoke(app, ["--version"]) assert result.exit_code == 0 - assert f"discolike-cli {version('discolike-cli')}" in result.output - assert f"(discolike {__version__})" in result.output + assert result.output.strip() == version("discolike-cli") + assert __version__ def test_status_honors_global_base_url( diff --git a/packages/discolike-cli/tests/test_auth_oauth.py b/packages/discolike-cli/tests/test_auth_oauth.py index 0ec8fdf..a0966c4 100644 --- a/packages/discolike-cli/tests/test_auth_oauth.py +++ b/packages/discolike-cli/tests/test_auth_oauth.py @@ -127,7 +127,7 @@ def test_login_default_runs_oauth_loopback_flow( stored = json.loads(config_path().read_text()) assert (stored["auth_method"], stored["oauth"]) == ("oauth", CREDENTIAL.to_config()) assert stored["oauth_client"] == {"client_id": "client-1", "redirect_uri": redirect_uri, "issuer": METADATA.issuer} - payload = json.loads(result.stderr.splitlines()[-1]) + payload = json.loads(result.stdout) assert payload == {"logged_in": True, "method": "oauth", "expires_at": "2027-01-15T08:00:00+00:00"} assert provider.opened_urls[0] in result.stderr diff --git a/packages/discolike-cli/tests/test_bulk_cli.py b/packages/discolike-cli/tests/test_bulk_cli.py new file mode 100644 index 0000000..5863988 --- /dev/null +++ b/packages/discolike-cli/tests/test_bulk_cli.py @@ -0,0 +1,660 @@ +"""discolike bulk companies|estimate|contacts (issue #24) — the pipeline mechanics, driven through a mock transport.""" + +from __future__ import annotations + +import csv +import json +from collections.abc import Callable +from pathlib import Path +from typing import Any + +import httpx2 +import pytest +from typer.testing import CliRunner + +from discolike_cli import bulk +from discolike_cli.main import app +from discolike_testkit import Handler +from discolike_testkit import plain_output + +runner = CliRunner() +NO_LIMIT = ["--rate-limit", "1000000"] + + +def _companies(prefix: str, count: int) -> list[dict[str, Any]]: + return [ + {"domain": f"{prefix}{i}.com", "name": f"{prefix.upper()} {i}", "employees": "11-50", "similarity": 0.9} + for i in range(count) + ] + + +def _fingerprint(domains: list[str], per_company: int) -> str: + """The checkpoint stamp for a contacts pull with only the default filters.""" + filters = bulk._contact_filters( + None, + None, + icp_prompt=None, + summary=None, + negate_summary=None, + seniority=None, + negate_seniority=None, + department=None, + negate_department=None, + title=None, + negate_title=None, + person_country=None, + person_state=None, + has_email=True, + exclusion_query_id=None, + ) + return f"fingerprint={bulk._checkpoint_fingerprint(domains, per_company, filters)}" + + +def _read_csv(path: Path) -> list[dict[str, str]]: + with path.open(newline="") as handle: + return list(csv.DictReader(handle)) + + +class Recorder: + """Records every request and answers from a per-path queue of responses.""" + + def __init__(self, **queues: list[Any]) -> None: + self.queues = queues + self.requests: list[httpx2.Request] = [] + + def __call__(self, request: httpx2.Request) -> httpx2.Response: + self.requests.append(request) + key = request.url.path.removeprefix("/v1/").replace("/", "_").replace("-", "_") + queue = self.queues[key] + payload = queue.pop(0) if len(queue) > 1 else queue[0] + if isinstance(payload, int): + return httpx2.Response(payload, json={"detail": "nope"}) + return httpx2.Response(200, json=payload) + + def bodies(self, path: str) -> list[dict[str, Any]]: + return [json.loads(r.content) for r in self.requests if r.url.path == path] + + def params(self, path: str) -> list[httpx2.QueryParams]: + return [r.url.params for r in self.requests if r.url.path == path] + + +@pytest.fixture(autouse=True) +def no_sleep(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr("time.sleep", lambda _seconds: None) + + +# --------------------------------------------------------------------------- companies + + +def test_bulk_companies_pages_with_exclusion_lists_and_appends_csv( + install_build_client: Callable[[Handler], None], tmp_path: Path +) -> None: + recorder = Recorder( + discover=[_companies("a", 20), _companies("b", 5)], + queries_exclusion_list=[{"query_id": "q-round-1"}], + ) + install_build_client(recorder) + out = tmp_path / "companies.csv" + result = runner.invoke( + app, + [ + "bulk", + "companies", + "--icp-prompt", + "agencies", + "--page-size", + "20", + "--max-companies", + "100", + "--run-name", + "agencies", + "--out", + str(out), + *NO_LIMIT, + ], + ) + assert result.exit_code == 0, result.output + + discover = recorder.params("/v1/discover") + assert [p.get("max_records") for p in discover] == ["20", "20"] + assert discover[0].get("exclusion_query_id") is None + assert discover[1].get_list("exclusion_query_id") == ["q-round-1"] + assert recorder.bodies("/v1/queries/exclusion-list") == [ + {"query_name": "agencies-round-1", "domains": [f"a{i}.com" for i in range(20)]} + ] + + rows = _read_csv(out) + assert len(rows) == 25 + assert rows[0] == {"domain": "a0.com", "name": "A 0", "country": "", "employees": "11-50", "similarity": "0.9"} + + summary = json.loads(result.stdout) + assert summary["companies"] == 25 + assert summary["rounds"] == 2 + assert summary["exclusion_query_ids"] == ["q-round-1"] + assert summary["out"] == str(out) + + +def test_bulk_companies_short_tail_rides_inline_exclude_domain( + install_build_client: Callable[[Handler], None], tmp_path: Path +) -> None: + page_one = _companies("a", 20) + page_two = [*_companies("a", 5), *_companies("b", 15)] # 5 already seen, 15 net-new (< 20, no saved list) + recorder = Recorder( + discover=[page_one, page_two, []], + queries_exclusion_list=[{"query_id": "q1"}], + ) + install_build_client(recorder) + result = runner.invoke( + app, + [ + "bulk", + "companies", + "--icp-prompt", + "x", + "--page-size", + "20", + "--max-companies", + "100", + "--out", + str(tmp_path / "c.csv"), + *NO_LIMIT, + ], + ) + assert result.exit_code == 0, result.output + assert len(recorder.bodies("/v1/queries/exclusion-list")) == 1 + third = recorder.params("/v1/discover")[2] + assert third.get_list("exclusion_query_id") == ["q1"] + assert third.get_list("exclude_domain") == [f"b{i}.com" for i in range(15)] + assert len(_read_csv(tmp_path / "c.csv")) == 35 + + +def test_bulk_companies_stops_at_max_companies(install_build_client: Callable[[Handler], None], tmp_path: Path) -> None: + recorder = Recorder( + discover=[_companies("a", 20), _companies("b", 20)], queries_exclusion_list=[{"query_id": "q1"}] + ) + install_build_client(recorder) + result = runner.invoke( + app, + [ + "bulk", + "companies", + "--icp-prompt", + "x", + "--page-size", + "20", + "--max-companies", + "40", + "--out", + str(tmp_path / "c.csv"), + *NO_LIMIT, + ], + ) + assert result.exit_code == 0, result.output + assert len(recorder.params("/v1/discover")) == 2 + assert len(_read_csv(tmp_path / "c.csv")) == 40 + + +def test_bulk_companies_resumes_from_existing_csv( + install_build_client: Callable[[Handler], None], tmp_path: Path +) -> None: + out = tmp_path / "companies.csv" + out.write_text("domain,name,country,employees,similarity\n" + "".join(f"old{i}.com,,,,\n" for i in range(25))) + recorder = Recorder(discover=[_companies("n", 3)], queries_exclusion_list=[{"query_id": "q-resume"}]) + install_build_client(recorder) + result = runner.invoke( + app, + [ + "bulk", + "companies", + "--icp-prompt", + "x", + "--page-size", + "20", + "--max-companies", + "100", + "--run-name", + "r", + "--out", + str(out), + *NO_LIMIT, + ], + ) + assert result.exit_code == 0, result.output + assert recorder.bodies("/v1/queries/exclusion-list") == [ + {"query_name": "r-resume-0", "domains": [f"old{i}.com" for i in range(25)]} + ] + assert recorder.params("/v1/discover")[0].get_list("exclusion_query_id") == ["q-resume"] + assert len(_read_csv(out)) == 28 + summary = json.loads(result.stdout) + assert summary == {**summary, "companies": 28, "new": 3} + + +def test_bulk_companies_resume_still_excludes_the_csv_when_a_suppression_list_is_supplied( + install_build_client: Callable[[Handler], None], tmp_path: Path +) -> None: + out = tmp_path / "companies.csv" + out.write_text("domain,name,country,employees,similarity\n" + "".join(f"old{i}.com,,,,\n" for i in range(25))) + recorder = Recorder(discover=[_companies("n", 3)], queries_exclusion_list=[{"query_id": "q-resume"}]) + install_build_client(recorder) + result = runner.invoke( + app, + [ + "bulk", + "companies", + "--icp-prompt", + "x", + "--page-size", + "20", + "--max-companies", + "100", + "--run-name", + "r", + "--exclusion-query-id", + "q-customers", + "--out", + str(out), + *NO_LIMIT, + ], + ) + assert result.exit_code == 0, result.output + assert recorder.bodies("/v1/queries/exclusion-list") == [ + {"query_name": "r-resume-0", "domains": [f"old{i}.com" for i in range(25)]} + ] + assert recorder.params("/v1/discover")[0].get_list("exclusion_query_id") == ["q-customers", "q-resume"] + + +def test_bulk_companies_overwrite_ignores_existing_csv( + install_build_client: Callable[[Handler], None], tmp_path: Path +) -> None: + out = tmp_path / "companies.csv" + out.write_text("domain,name,country,employees,similarity\nold.com,,,,\n") + recorder = Recorder(discover=[_companies("n", 3)], queries_exclusion_list=[{"query_id": "q"}]) + install_build_client(recorder) + result = runner.invoke(app, ["bulk", "companies", "--icp-prompt", "x", "--overwrite", "--out", str(out), *NO_LIMIT]) + assert result.exit_code == 0, result.output + assert recorder.bodies("/v1/queries/exclusion-list") == [] + assert [row["domain"] for row in _read_csv(out)] == ["n0.com", "n1.com", "n2.com"] + + +def test_bulk_companies_takes_params_file_and_drops_managed_keys( + install_build_client: Callable[[Handler], None], tmp_path: Path +) -> None: + form = tmp_path / "form.json" + form.write_text(json.dumps({"country": ["US"], "variance": "MID_HIGH", "max_records": 5, "offset": 99})) + recorder = Recorder(discover=[[]], queries_exclusion_list=[{"query_id": "q"}]) + install_build_client(recorder) + result = runner.invoke( + app, + [ + "bulk", + "companies", + "--params-file", + str(form), + "--param", + "consensus=3", + "--variance", + "LOW", + "--out", + str(tmp_path / "c.csv"), + *NO_LIMIT, + ], + ) + assert result.exit_code == 0, result.output + params = recorder.params("/v1/discover")[0] + assert params.get_list("country") == ["US"] + assert params.get("variance") == "LOW" + assert params.get("consensus") == "3" + assert params.get("max_records") == "10000" + assert params.get("offset") is None + assert "max_records" in result.output + assert "managed by bulk" in result.output + + +def test_bulk_companies_validates_before_first_call( + install_build_client: Callable[[Handler], None], tmp_path: Path +) -> None: + recorder = Recorder(discover=[[]]) + install_build_client(recorder) + result = runner.invoke( + app, ["bulk", "companies", "--variance", "BOGUS", "--out", str(tmp_path / "c.csv"), *NO_LIMIT] + ) + assert result.exit_code == 2 + assert "ValidationError" in result.output + assert recorder.requests == [] + + +def test_bulk_companies_retries_after_rate_limit( + install_build_client: Callable[[Handler], None], tmp_path: Path +) -> None: + # The SDK transport retries a 429 three times on its own; bulk retries the surfaced RateLimitError again. + recorder = Recorder(discover=[429, 429, 429, 429, _companies("a", 2)]) + install_build_client(recorder) + result = runner.invoke(app, ["bulk", "companies", "--icp-prompt", "x", "--out", str(tmp_path / "c.csv"), *NO_LIMIT]) + assert result.exit_code == 0, result.output + assert len(recorder.params("/v1/discover")) == 5 + assert len(_read_csv(tmp_path / "c.csv")) == 2 + + +# --------------------------------------------------------------------------- estimate + + +def test_bulk_estimate_counts_per_slice_and_caps( + install_build_client: Callable[[Handler], None], tmp_path: Path +) -> None: + domains = tmp_path / "companies.csv" + domains.write_text("domain\n" + "".join(f"d{i}.com\n" for i in range(2500))) + recorder = Recorder(contacts_count=[{"count": 30000}, {"count": 30000}, {"count": 1000}]) + install_build_client(recorder) + result = runner.invoke( + app, + ["bulk", "estimate", "--domains-file", str(domains), "--per-company", "10", "--seniority", "vp", *NO_LIMIT], + ) + assert result.exit_code == 0, result.output + calls = recorder.params("/v1/contacts/count") + assert [len(p.get_list("domain")) for p in calls] == [1000, 1000, 500] + assert calls[0].get_list("seniority") == ["vp"] + assert calls[0].get("has_email") == "true" + assert calls[0].get("max_records") is None + summary = json.loads(result.stdout) + assert summary == {"companies": 2500, "contacts_available": 61000, "contacts_capped": 25000, "per_company": 10} + + +# --------------------------------------------------------------------------- contacts + + +def _discover_payload(domains: list[str], per: int, start: int = 0) -> dict[str, Any]: + results = {} + for d_index, domain in enumerate(domains): + results[domain] = { + "domain": domain, + "name": domain.upper(), + "contacts": [ + { + "persona_id": start + d_index * per + c, + "name": f"First{c} Last{c}", + "title": "VP", + "email": f"p{c}@{domain}", + "phone": [{"phone": "+1"}], + "social_urls": [f"https://linkedin.com/in/p{c}"], + "industry": ["ACCOUNTING", "LEGAL"], + } + for c in range(per) + ], + } + return {"results": results} + + +def test_bulk_contacts_slices_domains_and_flattens_rows( + install_build_client: Callable[[Handler], None], tmp_path: Path +) -> None: + all_domains = [f"d{i}.com" for i in range(250)] + domains = tmp_path / "companies.csv" + domains.write_text("domain\n" + "".join(f"{d}\n" for d in all_domains)) + # per-company 100 => 10000 // 100 = 100 domains per call => slices of 100, 100, 50 + recorder = Recorder( + contacts_discover=[ + _discover_payload(all_domains[:2], 2, start=0), + _discover_payload(all_domains[100:102], 2, start=100), + _discover_payload(all_domains[200:201], 2, start=200), + ] + ) + install_build_client(recorder) + out = tmp_path / "contacts.csv" + result = runner.invoke( + app, + [ + "bulk", + "contacts", + "--domains-file", + str(domains), + "--per-company", + "100", + "--summary", + "growth", + "--negate-summary", + "bookkeeping", + "--out", + str(out), + *NO_LIMIT, + ], + ) + assert result.exit_code == 0, result.output + bodies = recorder.bodies("/v1/contacts/discover") + assert [b["domain"] for b in bodies] == [all_domains[:100], all_domains[100:200], all_domains[200:]] + assert bodies[0]["results_by_company"] == 100 + assert bodies[0]["max_records"] == 10000 + assert bodies[2]["max_records"] == 5000 + assert bodies[0]["summary"] == "growth" + stamp, *indexes = (tmp_path / "contacts.csv.checkpoint").read_text().split() + assert stamp.startswith("fingerprint=") + assert indexes == ["0", "1", "2"] + assert bodies[0]["negate_summary"] == "bookkeeping" + assert bodies[0]["has_email"] is True + assert "offset" not in bodies[0] + + rows = _read_csv(out) + assert len(rows) == 10 + assert rows[0]["persona_id"] == "0" + assert rows[0]["first_name"] == "First0" + assert rows[0]["last_name"] == "Last0" + assert rows[0]["email"] == "p0@d0.com" + assert rows[0]["phone"] == "+1" + assert rows[0]["linkedin"] == "https://linkedin.com/in/p0" + assert rows[0]["industry"] == "ACCOUNTING;LEGAL" + assert rows[0]["domain"] == "d0.com" + assert rows[0]["company_name"] == "D0.COM" + + summary = json.loads(result.stdout) + assert summary["companies"] == 250 + assert summary["slices"] == 3 + assert summary["contacts"] == 10 + assert summary["out"] == str(out) + + +def test_bulk_contacts_resume_skips_checkpointed_slices( + install_build_client: Callable[[Handler], None], tmp_path: Path +) -> None: + all_domains = [f"d{i}.com" for i in range(250)] + domains = tmp_path / "companies.csv" + domains.write_text("domain\n" + "".join(f"{d}\n" for d in all_domains)) + out = tmp_path / "contacts.csv" + out.write_text("persona_id,domain\n1,a.com\n") + (tmp_path / "contacts.csv.checkpoint").write_text(f"{_fingerprint(all_domains, 100)}\n0\n1\n") + recorder = Recorder(contacts_discover=[_discover_payload(all_domains[200:201], 1, start=200)]) + install_build_client(recorder) + result = runner.invoke( + app, + ["bulk", "contacts", "--domains-file", str(domains), "--per-company", "100", "--out", str(out), *NO_LIMIT], + ) + assert result.exit_code == 0, result.output + assert [b["domain"] for b in recorder.bodies("/v1/contacts/discover")] == [all_domains[200:]] + assert (tmp_path / "contacts.csv.checkpoint").read_text().split() == [_fingerprint(all_domains, 100), "0", "1", "2"] + assert out.read_text().splitlines()[1] == "1,a.com" # existing rows kept, header not repeated + summary = json.loads(result.stdout) + assert summary["slices_run"] == 1 + + +def test_bulk_contacts_overwrite_clears_checkpoint( + install_build_client: Callable[[Handler], None], tmp_path: Path +) -> None: + domains = tmp_path / "companies.csv" + domains.write_text("domain\na.com\n") + out = tmp_path / "contacts.csv" + out.write_text("stale\n") + (tmp_path / "contacts.csv.checkpoint").write_text("0\n") + recorder = Recorder(contacts_discover=[_discover_payload(["a.com"], 1)]) + install_build_client(recorder) + result = runner.invoke( + app, ["bulk", "contacts", "--domains-file", str(domains), "--overwrite", "--out", str(out), *NO_LIMIT] + ) + assert result.exit_code == 0, result.output + assert len(recorder.bodies("/v1/contacts/discover")) == 1 + assert len(_read_csv(out)) == 1 + + +def test_bulk_contacts_dedupes_persona_ids_within_run( + install_build_client: Callable[[Handler], None], tmp_path: Path +) -> None: + domains = tmp_path / "companies.csv" + domains.write_text("domain\na.com\nb.com\n") + recorder = Recorder(contacts_discover=[_discover_payload(["a.com", "b.com"], 2, start=0)]) + install_build_client(recorder) + result = runner.invoke( + app, ["bulk", "contacts", "--domains-file", str(domains), "--out", str(tmp_path / "c.csv"), *NO_LIMIT] + ) + assert result.exit_code == 0, result.output + assert len(_read_csv(tmp_path / "c.csv")) == 4 + + +def test_bulk_contacts_drops_managed_keys_from_params_file( + install_build_client: Callable[[Handler], None], tmp_path: Path +) -> None: + domains = tmp_path / "companies.csv" + domains.write_text("domain\na.com\n") + form = tmp_path / "form.json" + form.write_text(json.dumps({"domain": ["z.com"], "results_by_company": 1, "offset": 5, "seniority": ["vp"]})) + recorder = Recorder(contacts_discover=[_discover_payload(["a.com"], 1)]) + install_build_client(recorder) + result = runner.invoke( + app, + [ + "bulk", + "contacts", + "--domains-file", + str(domains), + "--params-file", + str(form), + "--per-company", + "3", + "--out", + str(tmp_path / "c.csv"), + *NO_LIMIT, + ], + ) + assert result.exit_code == 0, result.output + body = recorder.bodies("/v1/contacts/discover")[0] + assert body["domain"] == ["a.com"] + assert body["results_by_company"] == 3 + assert "offset" not in body + assert body["seniority"] == ["vp"] + assert "managed by bulk" in result.output + + +def test_bulk_companies_trims_last_page_to_max_companies( + install_build_client: Callable[[Handler], None], tmp_path: Path +) -> None: + # The API floor is 20 per call, so the final page can return more than the run still needs. + recorder = Recorder(discover=[_companies("a", 20), _companies("b", 20)], queries_exclusion_list=[{"query_id": "q"}]) + install_build_client(recorder) + result = runner.invoke( + app, + ["bulk", "companies", "--icp-prompt", "x", "--page-size", "20", "--max-companies", "30", + "--out", str(tmp_path / "c.csv"), *NO_LIMIT], + ) # fmt: skip + assert result.exit_code == 0, result.output + assert len(_read_csv(tmp_path / "c.csv")) == 30 + assert json.loads(result.stdout)["companies"] == 30 + + +def test_bulk_contacts_resume_dedupes_against_rows_already_in_csv( + install_build_client: Callable[[Handler], None], tmp_path: Path +) -> None: + domains = tmp_path / "companies.csv" + domains.write_text("domain\na.com\n") + out = tmp_path / "contacts.csv" + out.write_text("persona_id,domain\n0,a.com\n") # a lost checkpoint reruns slice 0 + recorder = Recorder(contacts_discover=[_discover_payload(["a.com"], 2)]) + install_build_client(recorder) + result = runner.invoke(app, ["bulk", "contacts", "--domains-file", str(domains), "--out", str(out), *NO_LIMIT]) + assert result.exit_code == 0, result.output + assert [row["persona_id"] for row in _read_csv(out)] == ["0", "1"] + + +def test_bulk_contacts_refuses_checkpoint_written_for_other_inputs( + install_build_client: Callable[[Handler], None], tmp_path: Path +) -> None: + domains = tmp_path / "companies.csv" + domains.write_text("domain\na.com\nb.com\n") + out = tmp_path / "contacts.csv" + out.write_text("persona_id,domain\n0,a.com\n") + (tmp_path / "contacts.csv.checkpoint").write_text(f"{_fingerprint(['a.com'], 100)}\n0\n") # b.com added since + recorder = Recorder(contacts_discover=[_discover_payload(["a.com", "b.com"], 1)]) + install_build_client(recorder) + result = runner.invoke(app, ["bulk", "contacts", "--domains-file", str(domains), "--out", str(out), *NO_LIMIT]) + assert result.exit_code == 2, result.output + assert "written for a different domains file" in plain_output(result.output) + assert recorder.requests == [] # nothing billed, nothing skipped + # --per-company changes the slicing, so the same domains still refuse to resume + (tmp_path / "contacts.csv.checkpoint").write_text(f"{_fingerprint(['a.com', 'b.com'], 50)}\n0\n") + result = runner.invoke(app, ["bulk", "contacts", "--domains-file", str(domains), "--out", str(out), *NO_LIMIT]) + assert result.exit_code == 2, result.output + # --overwrite discards it + result = runner.invoke( + app, ["bulk", "contacts", "--domains-file", str(domains), "--out", str(out), "--overwrite", *NO_LIMIT] + ) + assert result.exit_code == 0, result.output + assert len(_read_csv(out)) == 2 + + +def test_bulk_contacts_writes_every_contact_without_persona_id( + install_build_client: Callable[[Handler], None], tmp_path: Path +) -> None: + domains = tmp_path / "companies.csv" + domains.write_text("domain\na.com\n") + payload = _discover_payload(["a.com"], 3) + for contact in payload["results"]["a.com"]["contacts"]: + contact["persona_id"] = None + recorder = Recorder(contacts_discover=[payload]) + install_build_client(recorder) + result = runner.invoke( + app, ["bulk", "contacts", "--domains-file", str(domains), "--out", str(tmp_path / "c.csv"), *NO_LIMIT] + ) + assert result.exit_code == 0, result.output + assert [row["email"] for row in _read_csv(tmp_path / "c.csv")] == ["p0@a.com", "p1@a.com", "p2@a.com"] + + +def test_bulk_companies_consolidates_inline_tails_into_a_saved_list( + install_build_client: Callable[[Handler], None], tmp_path: Path +) -> None: + # Every page is 20 companies of which only 10 are net-new: each tail is too short for a saved list on + # its own, so tails ride on exclude_domain until two of them together reach the 20-domain minimum. + pages = [ + _companies("a", 20), + [*_companies("a", 10), *_companies("b", 10)], + [*_companies("b", 10), *_companies("c", 10)], + [*_companies("c", 10), *_companies("d", 10)], + [], + ] + recorder = Recorder(discover=pages, queries_exclusion_list=[{"query_id": "q1"}, {"query_id": "q2"}]) + install_build_client(recorder) + result = runner.invoke( + app, + ["bulk", "companies", "--icp-prompt", "x", "--page-size", "20", "--out", str(tmp_path / "c.csv"), *NO_LIMIT], + ) + assert result.exit_code == 0, result.output + lists = recorder.bodies("/v1/queries/exclusion-list") + assert [sorted(body["domains"]) for body in lists] == [ + sorted(c["domain"] for c in _companies("a", 20)), + sorted(c["domain"] for c in [*_companies("b", 10), *_companies("c", 10)]), + ] + assert lists[1]["query_name"] == "bulk-round-3-tails" + calls = recorder.params("/v1/discover") + assert calls[2].get_list("exclude_domain") == [c["domain"] for c in _companies("b", 10)] + assert calls[3].get_list("exclusion_query_id") == ["q1", "q2"] + assert calls[3].get_list("exclude_domain") == [] # consolidated, nothing inline + assert calls[4].get_list("exclude_domain") == [c["domain"] for c in _companies("d", 10)] + assert len(_read_csv(tmp_path / "c.csv")) == 50 + + +@pytest.mark.parametrize("command", [["companies", "--icp-prompt", "x"], ["estimate"], ["contacts"]]) +def test_bulk_rejects_non_positive_rate_limit( + install_build_client: Callable[[Handler], None], tmp_path: Path, command: list[str] +) -> None: + domains = tmp_path / "companies.csv" + domains.write_text("domain\na.com\n") + recorder = Recorder() + install_build_client(recorder) + extra = [] if command[0] == "companies" else ["--domains-file", str(domains)] + result = runner.invoke(app, ["bulk", *command, *extra, "--out", str(tmp_path / "c.csv"), "--rate-limit", "0"]) + assert result.exit_code == 2, result.output + assert recorder.requests == [] diff --git a/packages/discolike-cli/tests/test_contacts_cli.py b/packages/discolike-cli/tests/test_contacts_cli.py index d835bb8..002ede3 100644 --- a/packages/discolike-cli/tests/test_contacts_cli.py +++ b/packages/discolike-cli/tests/test_contacts_cli.py @@ -499,3 +499,19 @@ def handler(request: httpx2.Request) -> httpx2.Response: "full_domains": ["done.com"], "partial_domains": ["half.com"], } + + +def test_contacts_generate_find_emails_flag(install_build_client: Callable[[Handler], None]) -> None: + captured: dict[str, object] = {} + + def handler(request: httpx2.Request) -> httpx2.Response: + captured["body"] = json.loads(request.content) + return httpx2.Response(200, json={"task_id": "dg-9"}) + + install_build_client(handler) + result = runner.invoke( + app, + ["contacts", "generate", "--icp-text", "VPs of Marketing", "--domain", "acme.com", "--find-emails"], + ) + assert result.exit_code == 0, result.output + assert captured["body"] == {"icp_text": "VPs of Marketing", "domains": ["acme.com"], "find_emails": True} diff --git a/packages/discolike-cli/tests/test_discogen_cli.py b/packages/discolike-cli/tests/test_discogen_cli.py index 4397665..7095e79 100644 --- a/packages/discolike-cli/tests/test_discogen_cli.py +++ b/packages/discolike-cli/tests/test_discogen_cli.py @@ -67,6 +67,91 @@ def handler(request: httpx2.Request) -> httpx2.Response: } +def test_discogen_run_sends_typed_columns_when_passed(install_build_client: Callable[[Handler], None]) -> None: + captured: dict[str, object] = {} + + def handler(request: httpx2.Request) -> httpx2.Response: + captured["body"] = json.loads(request.content) + return httpx2.Response(200, json={"task_id": "dg-1c"}) + + install_build_client(handler) + result = runner.invoke( + app, + ["discogen", "run", "--query", "q", "--domain", "acme.com", "--typed-columns"], + ) + assert result.exit_code == 0, result.output + assert captured["body"] == { + "query": "q", + "domains": ["acme.com"], + "typed_columns": True, + } + + +def test_discogen_run_sends_include_confidence_true_when_flag_passed( + install_build_client: Callable[[Handler], None], +) -> None: + captured: dict[str, object] = {} + + def handler(request: httpx2.Request) -> httpx2.Response: + captured["body"] = json.loads(request.content) + return httpx2.Response(200, json={"task_id": "dg-1d"}) + + install_build_client(handler) + result = runner.invoke( + app, + ["discogen", "run", "--query", "q", "--domain", "acme.com", "--include-confidence"], + ) + assert result.exit_code == 0, result.output + assert captured["body"] == { + "query": "q", + "domains": ["acme.com"], + "include_confidence": True, + } + + +def test_discogen_run_sends_include_confidence_false_when_flag_negated( + install_build_client: Callable[[Handler], None], +) -> None: + captured: dict[str, object] = {} + + def handler(request: httpx2.Request) -> httpx2.Response: + captured["body"] = json.loads(request.content) + return httpx2.Response(200, json={"task_id": "dg-1e"}) + + install_build_client(handler) + result = runner.invoke( + app, + ["discogen", "run", "--query", "q", "--domain", "acme.com", "--no-include-confidence"], + ) + assert result.exit_code == 0, result.output + assert captured["body"] == { + "query": "q", + "domains": ["acme.com"], + "include_confidence": False, + } + + +def test_discogen_run_omits_include_confidence_when_flag_not_passed( + install_build_client: Callable[[Handler], None], +) -> None: + captured: dict[str, object] = {} + + def handler(request: httpx2.Request) -> httpx2.Response: + captured["body"] = json.loads(request.content) + return httpx2.Response(200, json={"task_id": "dg-1f"}) + + install_build_client(handler) + result = runner.invoke( + app, + ["discogen", "run", "--query", "q", "--domain", "acme.com"], + ) + assert result.exit_code == 0, result.output + assert captured["body"] == { + "query": "q", + "domains": ["acme.com"], + } + + def test_discogen_run_personas_posts_persona_ids(install_build_client: Callable[[Handler], None]) -> None: captured: dict[str, object] = {} @@ -110,6 +195,93 @@ def handler(request: httpx2.Request) -> httpx2.Response: } +def test_discogen_run_personas_sends_typed_columns_when_passed( + install_build_client: Callable[[Handler], None], +) -> None: + captured: dict[str, object] = {} + + def handler(request: httpx2.Request) -> httpx2.Response: + captured["body"] = json.loads(request.content) + return httpx2.Response(200, json={"task_id": "dg-2c"}) + + install_build_client(handler) + result = runner.invoke( + app, + ["discogen", "run-personas", "--query", "q", "--persona-id", "1", "--typed-columns"], + ) + assert result.exit_code == 0, result.output + assert captured["body"] == { + "query": "q", + "persona_ids": [1], + "typed_columns": True, + } + + +def test_discogen_run_personas_sends_include_confidence_true_when_flag_passed( + install_build_client: Callable[[Handler], None], +) -> None: + captured: dict[str, object] = {} + + def handler(request: httpx2.Request) -> httpx2.Response: + captured["body"] = json.loads(request.content) + return httpx2.Response(200, json={"task_id": "dg-2c"}) + + install_build_client(handler) + result = runner.invoke( + app, + ["discogen", "run-personas", "--query", "q", "--persona-id", "1", "--include-confidence"], + ) + assert result.exit_code == 0, result.output + assert captured["body"] == { + "query": "q", + "persona_ids": [1], + "include_confidence": True, + } + + +def test_discogen_run_personas_sends_include_confidence_false_when_flag_negated( + install_build_client: Callable[[Handler], None], +) -> None: + captured: dict[str, object] = {} + + def handler(request: httpx2.Request) -> httpx2.Response: + captured["body"] = json.loads(request.content) + return httpx2.Response(200, json={"task_id": "dg-2d"}) + + install_build_client(handler) + result = runner.invoke( + app, + ["discogen", "run-personas", "--query", "q", "--persona-id", "1", "--no-include-confidence"], + ) + assert result.exit_code == 0, result.output + assert captured["body"] == { + "query": "q", + "persona_ids": [1], + "include_confidence": False, + } + + +def test_discogen_run_personas_omits_include_confidence_when_flag_not_passed( + install_build_client: Callable[[Handler], None], +) -> None: + captured: dict[str, object] = {} + + def handler(request: httpx2.Request) -> httpx2.Response: + captured["body"] = json.loads(request.content) + return httpx2.Response(200, json={"task_id": "dg-2e"}) + + install_build_client(handler) + result = runner.invoke( + app, + ["discogen", "run-personas", "--query", "q", "--persona-id", "1"], + ) + assert result.exit_code == 0, result.output + assert captured["body"] == { + "query": "q", + "persona_ids": [1], + } + + def test_discogen_run_with_wait_polls_to_completion(install_build_client: Callable[[Handler], None]) -> None: statuses = iter( [ diff --git a/packages/discolike-cli/tests/test_discover.py b/packages/discolike-cli/tests/test_discover.py index edf001e..477cebc 100644 --- a/packages/discolike-cli/tests/test_discover.py +++ b/packages/discolike-cli/tests/test_discover.py @@ -232,6 +232,41 @@ def handler(request: httpx2.Request) -> httpx2.Response: assert params.get("exclude_leadgen") == "true" +def test_count_forwards_sub_industry_and_geo_flags(install_build_client: Callable[[Handler], None]) -> None: + captured: dict[str, httpx2.QueryParams] = {} + + def handler(request: httpx2.Request) -> httpx2.Response: + captured["params"] = request.url.params + return _count_ok(request) + + install_build_client(handler) + result = runner.invoke( + app, + [ + "count", + "--sub-industry", + "ROOFING", + "--negate-sub-industry", + "FOUNDRIES", + "--negate-sub-industry", + "SCAFFOLDING", + "--lat", + "40.7128", + "--lon", + "-74.006", + "--radius", + "30mi", + ], + ) + assert result.exit_code == 0, result.output + params = captured["params"] + assert params.get_list("sub_industry") == ["ROOFING"] + assert params.get_list("negate_sub_industry") == ["FOUNDRIES", "SCAFFOLDING"] + assert params.get("lat") == "40.7128" + assert params.get("lon") == "-74.006" + assert params.get("radius") == "30mi" + + def test_count_sends_shared_filter_subset(install_build_client: Callable[[Handler], None]) -> None: captured: dict[str, httpx2.QueryParams] = {} @@ -253,6 +288,120 @@ def test_count_param_without_equals_exits_2(install_build_client: Callable[[Hand assert result.exit_code == 2 +def test_discover_forwards_sub_industry_and_geo_flags(install_build_client: Callable[[Handler], None]) -> None: + captured: dict[str, httpx2.QueryParams] = {} + + def handler(request: httpx2.Request) -> httpx2.Response: + captured["params"] = request.url.params + return _discover_ok(request) + + install_build_client(handler) + result = runner.invoke( + app, + [ + "discover", + "--sub-industry", + "ROOFING", + "--sub-industry", + "CONSTRUCTION/ROOFING", + "--negate-sub-industry", + "FOUNDRIES", + "--lat", + "40.7128", + "--lon", + "-74.006", + "--radius", + "30mi", + ], + ) + assert result.exit_code == 0, result.output + params = captured["params"] + assert params.get_list("sub_industry") == ["ROOFING", "CONSTRUCTION/ROOFING"] + assert params.get_list("negate_sub_industry") == ["FOUNDRIES"] + assert params.get("lat") == "40.7128" + assert params.get("lon") == "-74.006" + assert params.get("radius") == "30mi" + + +def test_count_forwards_bbox(install_build_client: Callable[[Handler], None]) -> None: + captured: dict[str, httpx2.QueryParams] = {} + + def handler(request: httpx2.Request) -> httpx2.Response: + captured["params"] = request.url.params + return _count_ok(request) + + install_build_client(handler) + result = runner.invoke(app, ["count", "--bbox", "40.4,-74.3,41.0,-73.7"]) + assert result.exit_code == 0, result.output + assert captured["params"].get("bbox") == "40.4,-74.3,41.0,-73.7" + + +def test_discover_forwards_bbox(install_build_client: Callable[[Handler], None]) -> None: + captured: dict[str, httpx2.QueryParams] = {} + + def handler(request: httpx2.Request) -> httpx2.Response: + captured["params"] = request.url.params + return _discover_ok(request) + + install_build_client(handler) + result = runner.invoke(app, ["discover", "--bbox", "40.4,-74.3,41.0,-73.7"]) + assert result.exit_code == 0, result.output + assert captured["params"].get("bbox") == "40.4,-74.3,41.0,-73.7" + + +def test_discover_forwards_several_geo_shapes(install_build_client: Callable[[Handler], None]) -> None: + captured: dict[str, httpx2.QueryParams] = {} + + def handler(request: httpx2.Request) -> httpx2.Response: + captured["params"] = request.url.params + return _discover_ok(request) + + install_build_client(handler) + result = runner.invoke( + app, + [ + "discover", + "--geo", + "30.27,-97.74,10km", + "--geo", + "52.52,13.405", + "--bbox", + "40.4,-74.3,41.0,-73.7", + "--bbox", + "51.2,-0.5,51.7,0.3", + ], + ) + assert result.exit_code == 0, result.output + params = captured["params"] + assert params.get_list("geo") == ["30.27,-97.74,10km", "52.52,13.405"] + assert params.get_list("bbox") == ["40.4,-74.3,41.0,-73.7", "51.2,-0.5,51.7,0.3"] + + +def test_count_forwards_several_geo_shapes(install_build_client: Callable[[Handler], None]) -> None: + captured: dict[str, httpx2.QueryParams] = {} + + def handler(request: httpx2.Request) -> httpx2.Response: + captured["params"] = request.url.params + return _count_ok(request) + + install_build_client(handler) + result = runner.invoke(app, ["count", "--geo", "30.27,-97.74", "--geo", "52.52,13.405,30mi"]) + assert result.exit_code == 0, result.output + assert captured["params"].get_list("geo") == ["30.27,-97.74", "52.52,13.405,30mi"] + + +def test_discover_rejects_a_bad_geo_circle(install_build_client: Callable[[Handler], None]) -> None: + install_build_client(_discover_ok) + result = runner.invoke(app, ["discover", "--geo", "30.27,-97.74,0km"]) + assert result.exit_code != 0 + + +def test_discover_rejects_out_of_range_bbox(install_build_client: Callable[[Handler], None]) -> None: + install_build_client(_discover_ok) + result = runner.invoke(app, ["discover", "--bbox", "40.4,-74.3,91.0,-73.7"]) + assert result.exit_code != 0 + + def test_discover_unauthorized_exits_3(install_build_client: Callable[[Handler], None]) -> None: install_build_client(_unauthorized) result = runner.invoke(app, ["discover", "--icp-prompt", "X"]) diff --git a/packages/discolike-cli/tests/test_enrich_cli.py b/packages/discolike-cli/tests/test_enrich_cli.py index 1cb3a3c..dccec9a 100644 --- a/packages/discolike-cli/tests/test_enrich_cli.py +++ b/packages/discolike-cli/tests/test_enrich_cli.py @@ -271,3 +271,23 @@ def handler(request: httpx2.Request) -> httpx2.Response: assert result.exit_code == 0, result.output assert captured[0].url.params.get_list("query_id") == ["q1"] assert captured[1].url.params.get_list("query_id") == ["q2"] + + +def test_validate_icp_sends_native_icp_sentinel(install_build_client: Callable[[Handler], None]) -> None: + captured: dict[str, object] = {} + + def handler(request: httpx2.Request) -> httpx2.Response: + captured["body"] = json.loads(request.content) + return httpx2.Response(200, json={"task_id": "vi-native"}) + + install_build_client(handler) + result = runner.invoke( + app, + ["validate-icp", "--icp", "Cybersecurity for SMBs", "--domain", "acme.com", "--integration-id", "native-icp"], + ) + assert result.exit_code == 0, result.output + assert captured["body"] == { + "icp_text": "Cybersecurity for SMBs", + "domains": ["acme.com"], + "integration_id": "native-icp", + } diff --git a/packages/discolike-cli/tests/test_entrypoint.py b/packages/discolike-cli/tests/test_entrypoint.py new file mode 100644 index 0000000..452d96a --- /dev/null +++ b/packages/discolike-cli/tests/test_entrypoint.py @@ -0,0 +1,87 @@ +"""The console entry point wraps every parser failure in the JSON error envelope.""" + +from __future__ import annotations + +import json +import pathlib + +import pytest + +from discolike_cli.main import run + + +def _run(argv: list[str], capsys: pytest.CaptureFixture[str]) -> tuple[int, str, str]: + with pytest.raises(SystemExit) as raised: + run(argv) + captured = capsys.readouterr() + code = raised.value.code + return (code if isinstance(code, int) else 1), captured.out, captured.err + + +@pytest.mark.parametrize( + ("argv", "fragment"), + [ + (["count", "--nope"], "--nope"), + (["discover", "--max-records", "abc"], "not a valid integer"), + (["discogen", "status"], "TASK_ID"), + (["auth", "login", "--method", "bogus"], "--method"), + ], +) +def test_parser_errors_use_the_json_envelope( + argv: list[str], fragment: str, capsys: pytest.CaptureFixture[str] +) -> None: + code, out, err = _run(argv, capsys) + assert code == 2 + assert out == "" + payload = json.loads(err) + assert payload["code"] == "validation_error" + assert payload["error"] == "ValidationError" + assert payload["exit_code"] == 2 + assert fragment in payload["message"] + + +def test_help_still_renders(capsys: pytest.CaptureFixture[str]) -> None: + code, out, err = _run(["--help"], capsys) + assert code == 0 + assert "Exit codes" in out + assert err == "" + + +def test_subcommand_help_still_renders(capsys: pytest.CaptureFixture[str]) -> None: + code, out, _ = _run(["count", "--help"], capsys) + assert code == 0 + assert "Output (success" in out + + +def test_bare_invocation_prints_help_not_an_envelope(capsys: pytest.CaptureFixture[str]) -> None: + code, out, err = _run([], capsys) + assert code in (0, 2) + assert "Usage" in out + err + assert not err.startswith("{") + + +def test_version_exits_zero(capsys: pytest.CaptureFixture[str]) -> None: + code, out, _ = _run(["--version"], capsys) + assert code == 0 + assert out.strip() + + +def test_api_errors_keep_their_exit_code( + monkeypatch: pytest.MonkeyPatch, tmp_path: pathlib.Path, capsys: pytest.CaptureFixture[str] +) -> None: + monkeypatch.delenv("DISCOLIKE_API_KEY", raising=False) + monkeypatch.setenv("XDG_CONFIG_HOME", str(tmp_path)) # empty: no config file, no credential + code, _, err = _run(["count", "--country", "US"], capsys) + assert code == 3 + assert json.loads(err)["code"] == "auth_required" + + +def test_keyboard_interrupt_exits_130(monkeypatch: pytest.MonkeyPatch, capsys: pytest.CaptureFixture[str]) -> None: + import discolike_cli.main as cli_main + + def _boom(*_args: object, **_kwargs: object) -> None: + raise KeyboardInterrupt + + monkeypatch.setattr(cli_main, "app", _boom) + code, _, _ = _run(["count"], capsys) + assert code == 130 diff --git a/packages/discolike-cli/tests/test_file_inputs_cli.py b/packages/discolike-cli/tests/test_file_inputs_cli.py new file mode 100644 index 0000000..23d611c --- /dev/null +++ b/packages/discolike-cli/tests/test_file_inputs_cli.py @@ -0,0 +1,210 @@ +"""--domains-file / --params-file / --exclude-domains-file on the volume commands (issue #23).""" + +from __future__ import annotations + +import json +from collections.abc import Callable +from pathlib import Path +from typing import Any + +import httpx2 +import pytest +from typer.testing import CliRunner + +from discolike_cli.main import app +from discolike_testkit import Handler +from discolike_testkit import plain_output + +runner = CliRunner() + + +@pytest.fixture +def domains_file(tmp_path: Path) -> Path: + path = tmp_path / "companies.csv" + path.write_text("domain,name\nacme.com,Acme\nbeta.io,Beta\n") + return path + + +def _capture_json(captured: dict[str, Any], response: object) -> Handler: + def handler(request: httpx2.Request) -> httpx2.Response: + captured["path"] = request.url.path + captured["params"] = request.url.params + captured["body"] = json.loads(request.content) if request.content else None + return httpx2.Response(200, json=response) + + return handler + + +def test_create_exclusion_list_merges_domains_file_with_inline( + install_build_client: Callable[[Handler], None], domains_file: Path +) -> None: + captured: dict[str, Any] = {} + install_build_client(_capture_json(captured, {"query_id": "q1"})) + result = runner.invoke( + app, + [ + "queries", + "create-exclusion-list", + "--name", + "L", + "--domain", + "gamma.co", + "--domains-file", + str(domains_file), + ], + ) + assert result.exit_code == 0, result.output + assert captured["body"] == {"query_name": "L", "domains": ["gamma.co", "acme.com", "beta.io"]} + + +def test_contacts_discover_domains_file(install_build_client: Callable[[Handler], None], domains_file: Path) -> None: + captured: dict[str, Any] = {} + install_build_client(_capture_json(captured, {"results": {}})) + result = runner.invoke(app, ["contacts", "discover", "--domains-file", str(domains_file)]) + assert result.exit_code == 0, result.output + assert captured["body"] == {"domain": ["acme.com", "beta.io"]} + + +def test_contacts_search_domains_file(install_build_client: Callable[[Handler], None], domains_file: Path) -> None: + captured: dict[str, Any] = {} + install_build_client(_capture_json(captured, [])) + result = runner.invoke(app, ["contacts", "search", "--domains-file", str(domains_file)]) + assert result.exit_code == 0, result.output + assert captured["params"].get_list("domain") == ["acme.com", "beta.io"] + + +def test_contacts_count_domains_file(install_build_client: Callable[[Handler], None], domains_file: Path) -> None: + captured: dict[str, Any] = {} + install_build_client(_capture_json(captured, {"count": 3})) + result = runner.invoke(app, ["contacts", "count", "--domains-file", str(domains_file)]) + assert result.exit_code == 0, result.output + assert captured["params"].get_list("domain") == ["acme.com", "beta.io"] + + +def test_contacts_generate_domains_file_replaces_required_domain( + install_build_client: Callable[[Handler], None], domains_file: Path +) -> None: + captured: dict[str, Any] = {} + install_build_client(_capture_json(captured, {"task_id": "t1"})) + result = runner.invoke(app, ["contacts", "generate", "--icp-text", "X", "--domains-file", str(domains_file)]) + assert result.exit_code == 0, result.output + assert captured["body"] == {"icp_text": "X", "domains": ["acme.com", "beta.io"]} + + +def test_contacts_generate_requires_domain_or_file(install_build_client: Callable[[Handler], None]) -> None: + install_build_client(_capture_json({}, {"task_id": "t1"})) + result = runner.invoke(app, ["contacts", "generate", "--icp-text", "X"]) + assert result.exit_code != 0 + assert "--domain or --domains-file" in plain_output(result.output) + + +def test_discogen_run_domains_file(install_build_client: Callable[[Handler], None], domains_file: Path) -> None: + captured: dict[str, Any] = {} + install_build_client(_capture_json(captured, {"task_id": "t1"})) + result = runner.invoke(app, ["discogen", "run", "--query", "Q", "--domains-file", str(domains_file)]) + assert result.exit_code == 0, result.output + assert captured["body"] == {"query": "Q", "domains": ["acme.com", "beta.io"]} + + +def test_validate_icp_accepts_domains_file_alias( + install_build_client: Callable[[Handler], None], domains_file: Path +) -> None: + captured: dict[str, Any] = {} + install_build_client(_capture_json(captured, {"task_id": "t1"})) + result = runner.invoke(app, ["validate-icp", "--icp", "X", "--domains-file", str(domains_file)]) + assert result.exit_code == 0, result.output + assert captured["body"] == {"icp_text": "X", "domains": ["acme.com", "beta.io"]} + + +def test_discover_params_file_under_param_under_flags( + install_build_client: Callable[[Handler], None], tmp_path: Path +) -> None: + form = tmp_path / "form.json" + form.write_text(json.dumps({"country": ["US"], "variance": "LOW", "min_similarity": 10, "consensus": 3})) + captured: dict[str, Any] = {} + install_build_client(_capture_json(captured, [])) + result = runner.invoke( + app, + ["discover", "--params-file", str(form), "--param", "min_similarity=50", "--variance", "HIGH"], + ) + assert result.exit_code == 0, result.output + params = captured["params"] + assert params.get_list("country") == ["US"] + assert params.get("variance") == "HIGH" + assert params.get("min_similarity") == "50" + assert params.get("consensus") == "3" + + +def test_discover_params_file_is_validated_locally( + install_build_client: Callable[[Handler], None], tmp_path: Path +) -> None: + form = tmp_path / "form.json" + form.write_text(json.dumps({"variance": "BOGUS"})) + calls: list[str] = [] + + def handler(request: httpx2.Request) -> httpx2.Response: + calls.append(request.url.path) + return httpx2.Response(200, json=[]) + + install_build_client(handler) + result = runner.invoke(app, ["discover", "--params-file", str(form)]) + assert result.exit_code == 2 + assert "ValidationError" in result.output + assert calls == [] + + +def test_count_params_file(install_build_client: Callable[[Handler], None], tmp_path: Path) -> None: + form = tmp_path / "form.json" + form.write_text(json.dumps({"country": ["DE"]})) + captured: dict[str, Any] = {} + install_build_client(_capture_json(captured, {"count": 1})) + result = runner.invoke(app, ["count", "--params-file", str(form)]) + assert result.exit_code == 0, result.output + assert captured["params"].get_list("country") == ["DE"] + + +@pytest.mark.parametrize("command", ["discover", "search", "count"]) +def test_contacts_params_file(install_build_client: Callable[[Handler], None], tmp_path: Path, command: str) -> None: + form = tmp_path / "form.json" + form.write_text(json.dumps({"seniority": ["executive"], "negate_summary": "bookkeeping"})) + captured: dict[str, Any] = {} + response: object = {"results": {}} if command == "discover" else ([] if command == "search" else {"count": 0}) + install_build_client(_capture_json(captured, response)) + result = runner.invoke(app, ["contacts", command, "--params-file", str(form), "--param", "seniority=vp"]) + assert result.exit_code == 0, result.output + if command == "discover": + assert captured["body"] == {"seniority": ["vp"], "negate_summary": "bookkeeping"} + else: + assert captured["params"].get_list("seniority") == ["vp"] + assert captured["params"].get("negate_summary") == "bookkeeping" + + +def test_discover_exclude_domains_file_merges_with_inline( + install_build_client: Callable[[Handler], None], domains_file: Path +) -> None: + captured: dict[str, Any] = {} + install_build_client(_capture_json(captured, [])) + result = runner.invoke( + app, + ["discover", "--icp-prompt", "X", "--exclude-domain", "gamma.co", "--exclude-domains-file", str(domains_file)], + ) + assert result.exit_code == 0, result.output + assert captured["params"].get_list("exclude_domain") == ["gamma.co", "acme.com", "beta.io"] + + +def test_discover_exclude_domains_file_over_cap_fails_before_call( + install_build_client: Callable[[Handler], None], tmp_path: Path +) -> None: + path = tmp_path / "many.txt" + path.write_text("\n".join(f"d{i}.com" for i in range(101))) + calls: list[str] = [] + + def handler(request: httpx2.Request) -> httpx2.Response: + calls.append(request.url.path) + return httpx2.Response(200, json=[]) + + install_build_client(handler) + result = runner.invoke(app, ["discover", "--icp-prompt", "X", "--exclude-domains-file", str(path)]) + assert result.exit_code == 2 + assert "ValidationError" in result.output + assert calls == [] diff --git a/packages/discolike-cli/tests/test_help.py b/packages/discolike-cli/tests/test_help.py new file mode 100644 index 0000000..d25bbfd --- /dev/null +++ b/packages/discolike-cli/tests/test_help.py @@ -0,0 +1,84 @@ +from __future__ import annotations + +import pytest +from typer.testing import CliRunner + +from discolike_cli import _help +from discolike_cli.main import app + +runner = CliRunner() + + +def test_top_level_help_documents_agent_contract() -> None: + result = runner.invoke(app, ["--help"]) + assert result.exit_code == 0 + assert "Output contract" in result.output + assert "Exit codes" in result.output + assert "auth_required" in result.output + assert "DISCOLIKE_API_KEY" in result.output + # The exit-code table must survive verbatim, one code per line, not reflowed into prose. + lines = [line.strip() for line in result.output.splitlines()] + assert any(line.startswith("2 validation_error") for line in lines) + assert any(line.startswith("6 not_found") for line in lines) + + +@pytest.mark.parametrize( + "argv", + [ + ["discover"], + ["count"], + ["match"], + ["append"], + ["validate-icp"], + ["segment"], + ["extract"], + ["signup"], + ["contacts", "search"], + ["contacts", "generate"], + ["discogen", "run"], + ["discogen", "status"], + ["queries", "create-exclusion-list"], + ["auth", "login"], + ["auth", "status"], + ["account", "usage"], + ], +) +def test_command_help_documents_output_and_errors(argv: list[str]) -> None: + result = runner.invoke(app, [*argv, "--help"]) + assert result.exit_code == 0, result.output + assert "Output (success" in result.output + assert "Common errors" in result.output + assert any(line.strip().startswith("Common errors:") for line in result.output.splitlines()) + + +def test_validate_icp_help_documents_a_title_case_verdict() -> None: + result = runner.invoke(app, ["validate-icp", "--help"]) + assert result.exit_code == 0, result.output + assert '"Fit": "Yes"|"No"' in result.output + assert '"ICP Fit" ("Yes"|"No"' in result.output + assert '"yes"' not in result.output + assert '"no"' not in result.output + + +def test_every_command_epilog_is_short() -> None: + for name, text in _help.COMMAND_EPILOGS.items(): + assert len(text.splitlines()) <= 18, name + + +def test_version_prints_bare_semver_when_not_a_tty() -> None: + from importlib.metadata import version + + result = runner.invoke(app, ["--version"]) + assert result.exit_code == 0 + assert result.output.strip() == version("discolike-cli") + + +def test_version_prints_decorated_line_on_a_tty(monkeypatch: pytest.MonkeyPatch) -> None: + from importlib.metadata import version + + import discolike_cli.main as cli_main + + monkeypatch.setattr(cli_main, "_stdout_is_tty", lambda: True) + result = runner.invoke(app, ["--version"]) + assert result.exit_code == 0 + assert f"discolike-cli {version('discolike-cli')}" in result.output diff --git a/packages/discolike-cli/tests/test_inputs.py b/packages/discolike-cli/tests/test_inputs.py new file mode 100644 index 0000000..0e14809 --- /dev/null +++ b/packages/discolike-cli/tests/test_inputs.py @@ -0,0 +1,81 @@ +from __future__ import annotations + +from pathlib import Path + +import pytest +import typer + +from discolike_cli._inputs import read_domains_file +from discolike_cli._inputs import read_params_file + + +def test_read_domains_file_one_per_line_normalizes_and_dedupes(tmp_path: Path) -> None: + path = tmp_path / "domains.txt" + path.write_text("Acme.com\nwww.beta.io\n\n acme.com \n") + assert read_domains_file(path) == ["acme.com", "beta.io"] + + +def test_read_domains_file_csv_uses_domain_column(tmp_path: Path) -> None: + path = tmp_path / "companies.csv" + path.write_text("name,Domain,country\nAcme,acme.com,US\nBeta,beta.io,DE\n") + assert read_domains_file(path) == ["acme.com", "beta.io"] + + +def test_read_domains_file_csv_without_domain_header_uses_first_column(tmp_path: Path) -> None: + path = tmp_path / "companies.csv" + path.write_text("acme.com,Acme\nbeta.io,Beta\n") + assert read_domains_file(path) == ["acme.com", "beta.io"] + + +def test_read_domains_file_missing_is_bad_parameter(tmp_path: Path) -> None: + with pytest.raises(typer.BadParameter, match="not found"): + read_domains_file(tmp_path / "nope.csv") + + +def test_read_domains_file_empty_is_bad_parameter(tmp_path: Path) -> None: + path = tmp_path / "empty.txt" + path.write_text("\n\n") + with pytest.raises(typer.BadParameter, match="no domains"): + read_domains_file(path) + + +def test_read_params_file_returns_object(tmp_path: Path) -> None: + path = tmp_path / "form.json" + path.write_text('{"country": ["US"], "variance": "MID_HIGH"}') + assert read_params_file(path) == {"country": ["US"], "variance": "MID_HIGH"} + + +def test_read_params_file_rejects_non_object(tmp_path: Path) -> None: + path = tmp_path / "form.json" + path.write_text("[1, 2]") + with pytest.raises(typer.BadParameter, match="JSON object"): + read_params_file(path) + + +def test_read_params_file_rejects_invalid_json(tmp_path: Path) -> None: + path = tmp_path / "form.json" + path.write_text("{nope") + with pytest.raises(typer.BadParameter, match="valid JSON"): + read_params_file(path) + + +def test_read_params_file_missing_is_bad_parameter(tmp_path: Path) -> None: + with pytest.raises(typer.BadParameter, match="not found"): + read_params_file(tmp_path / "nope.json") + + +def test_read_domains_file_directory_is_bad_parameter(tmp_path: Path) -> None: + with pytest.raises(typer.BadParameter, match="could not be read"): + read_domains_file(tmp_path) + + +def test_read_domains_file_bad_encoding_is_bad_parameter(tmp_path: Path) -> None: + path = tmp_path / "latin1.csv" + path.write_bytes(b"caf\xe9.com\n") + with pytest.raises(typer.BadParameter, match="could not be read"): + read_domains_file(path) + + +def test_read_params_file_directory_is_bad_parameter(tmp_path: Path) -> None: + with pytest.raises(typer.BadParameter, match="could not be read"): + read_params_file(tmp_path) diff --git a/packages/discolike-cli/tests/test_output.py b/packages/discolike-cli/tests/test_output.py index 106fa1c..fe6950f 100644 --- a/packages/discolike-cli/tests/test_output.py +++ b/packages/discolike-cli/tests/test_output.py @@ -8,13 +8,21 @@ import pytest import typer +from discolike import APIConnectionError from discolike import AuthenticationError +from discolike import DiscolikeError +from discolike import JobFailedError +from discolike import JobTimeoutError +from discolike import NotFoundError +from discolike import PlanAccessError from discolike import RateLimitError from discolike import ServerError from discolike import ValidationError +from discolike._config import NO_CREDENTIAL_MESSAGE from discolike.requests import DiscoverParams from discolike.requests import MatchCompanyParams from discolike.resources.discovery import Company +from discolike_cli._output import ERROR_CODES from discolike_cli._output import EXIT_CODES from discolike_cli._output import build_request from discolike_cli._output import emit @@ -160,7 +168,13 @@ def test_fail_writes_stderr_json_and_returns_typer_exit(capsys: pytest.CaptureFi assert result.exit_code == 2 captured = capsys.readouterr() payload = json.loads(captured.err) - assert payload == {"error": "ValidationError", "message": "bad field", "status_code": 400} + assert payload == { + "error": "ValidationError", + "code": "validation_error", + "message": "bad field", + "status_code": 400, + "exit_code": 2, + } assert captured.out == "" @@ -295,3 +309,35 @@ def bad() -> None: assert payload["status_code"] is None assert "name" in payload["message"] assert "min_match_confidence" in payload["message"] + + +@pytest.mark.parametrize( + ("exc", "code", "exit_code"), + [ + (ValidationError("bad", status_code=400), "validation_error", 2), + (AuthenticationError(NO_CREDENTIAL_MESSAGE), "auth_required", 3), + (AuthenticationError("OAuth token response has no `refresh_token`"), "auth_invalid", 3), + (AuthenticationError("OAuth session expired; run `discolike auth login`"), "auth_invalid", 3), + (AuthenticationError("bad key", status_code=401), "auth_invalid", 3), + (AuthenticationError("forbidden", status_code=403), "auth_invalid", 3), + (PlanAccessError("upgrade", status_code=403), "plan_access", 3), + (RateLimitError("slow", status_code=429, retry_after=1.0), "rate_limited", 4), + (APIConnectionError("down"), "network_error", 5), + (NotFoundError("nope", status_code=404), "not_found", 6), + (ServerError("boom", status_code=500), "server_error", 1), + (JobFailedError("failed"), "job_failed", 1), + (JobTimeoutError("slow job"), "job_timeout", 1), + ], +) +def test_fail_emits_stable_code_and_exit_code( + exc: DiscolikeError, code: str, exit_code: int, capsys: pytest.CaptureFixture[str] +) -> None: + result = fail(exc) + payload = json.loads(capsys.readouterr().err) + assert payload["code"] == code + assert payload["exit_code"] == exit_code + assert result.exit_code == exit_code + + +def test_error_codes_cover_every_exit_code_type() -> None: + assert set(EXIT_CODES) <= set(ERROR_CODES) diff --git a/packages/discolike-cli/tests/test_sdk_parity.py b/packages/discolike-cli/tests/test_sdk_parity.py index 49926f9..ba1b1c3 100644 --- a/packages/discolike-cli/tests/test_sdk_parity.py +++ b/packages/discolike-cli/tests/test_sdk_parity.py @@ -17,12 +17,18 @@ CLI_SOURCE_DIR = pathlib.Path(__file__).resolve().parents[1] / "src" / "discolike_cli" REQUEST_BUILDER = "build_request" # icp_text was dropped from discover and contacts in 0.1.1 in favor of icp_prompt; see the *_icp_text_*_removed tests. +# source_query_id/refs/round are saved-query plumbing: refs must be index-aligned with a stored +# query's contacts, which the CLI never holds, so there's no flag for them to bind to. DELIBERATELY_OMITTED: dict[str, frozenset[str]] = { "DiscoverParams": frozenset({"icp_text"}), "ContactsSearchParams": frozenset({"icp_text"}), "ContactsCountParams": frozenset({"icp_text"}), "ContactFilters": frozenset({"icp_text"}), + "FindEmailBatchRequest": frozenset({"source_query_id", "refs", "round"}), } +# ``discolike bulk`` takes the full vocabulary through --params-file / --param and manages the paging +# fields itself; its flags are the handful a volume run needs, not one per SDK field. +PARAMS_FILE_SITES = frozenset({"bulk.companies_command", "bulk.estimate_command", "bulk.contacts_command"}) DICT_ONLY_FIELDS: dict[str, frozenset[str]] = { "ContactGenerateRequest": frozenset({"initial_contact_counts"}), "DiscoGenProcessRequest": frozenset({"previous_discogen_data"}), @@ -77,7 +83,7 @@ def _build_request_sites() -> Iterator[tuple[str, str, set[str]]]: yield f"{source_file.stem}.{function.name}", call.args[0].id, forwarded -SITES = list(_build_request_sites()) +SITES = [site for site in _build_request_sites() if site[0] not in PARAMS_FILE_SITES] @pytest.mark.parametrize(("site", "model_name", "forwarded"), SITES, ids=[f"{s}->{m}" for s, m, _ in SITES]) diff --git a/packages/discolike-testkit/src/discolike_testkit/__init__.py b/packages/discolike-testkit/src/discolike_testkit/__init__.py index 8afa009..1c7ffbb 100644 --- a/packages/discolike-testkit/src/discolike_testkit/__init__.py +++ b/packages/discolike-testkit/src/discolike_testkit/__init__.py @@ -7,6 +7,7 @@ the calls made through them. """ +import re from collections.abc import Callable import httpx2 @@ -16,7 +17,7 @@ from discolike._auth import DiscolikeAuth from discolike._credentials import ApiKeyCredential -__all__ = ["AsyncClientFactory", "ClientFactory", "Handler", "api_key_auth"] +__all__ = ["AsyncClientFactory", "ClientFactory", "Handler", "api_key_auth", "plain_output"] Handler = Callable[[httpx2.Request], httpx2.Response] ClientFactory = Callable[[Handler], Discolike] @@ -26,3 +27,15 @@ def api_key_auth(api_key: str) -> DiscolikeAuth: """Auth for tests that build a ``Transport`` directly instead of going through ``Discolike``.""" return DiscolikeAuth(ApiKeyCredential(api_key=api_key)) + + +_ANSI = re.compile(r"\x1b\[[0-9;]*m") + + +def plain_output(output: str) -> str: + """Flatten CLI output for substring assertions. + + typer renders usage errors in a rich panel; on GitHub Actions rich forces color and wraps at + 80 columns, so a message that is one line locally arrives colored and split across box rows. + """ + return " ".join(_ANSI.sub("", output).replace("│", " ").split()) diff --git a/packages/discolike/README.md b/packages/discolike/README.md index db8ef62..152a58a 100644 --- a/packages/discolike/README.md +++ b/packages/discolike/README.md @@ -78,6 +78,50 @@ result = job.wait() `JobTimeoutError` is a client-side wait limit only — the task keeps running server-side (large DiscoGen runs can take hours), so call `wait()` again to resume or fetch `status()` later. Cancelled tasks still return results for every item that finished before cancellation. Send one job per list (up to 10,000 domains) rather than splitting into parallel jobs — concurrent DiscoGen jobs share your LLM provider key and slow each other down. +### Contact generation without an LLM key + +`contacts.generate` runs on your own search provider plus either your own LLM or DiscoLike Groove, DiscoLike's native extractor. Pass `NATIVE_ENGINE` to skip the LLM entirely: + +```python +from discolike import NATIVE_ENGINE +from discolike.requests import ContactGenerateRequest + +job = client.contacts.generate( + ContactGenerateRequest( + icp_text="VPs or Directors of Marketing at B2B SaaS", + domains=["gusto.com", "rippling.com"], + integration_id=NATIVE_ENGINE, + ) +) +result = job.wait() +print(result.title_validation) # "none" on the native engine, "llm" on a BYOK run +``` + +The native engine returns every person the search surfaces with a title and does not validate titles against `icp_text`, so filter them yourself when that matters. Omit `integration_id` to use your default LLM integration, or native when you have none. A search provider is required either way. + +### ICP validation without an LLM key + +`validate_icp` runs your ICP text against each domain on your own LLM provider key, or on DiscoLike's own ICP-fit model. Pass `NATIVE_ICP_ENGINE` for the latter — no LLM key, no LLM cost, and no web search on that run: + +```python +from discolike import NATIVE_ICP_ENGINE +from discolike.requests import ValidateIcpRequest + +job = client.validate_icp( + ValidateIcpRequest( + icp_text="Cybersecurity for SMBs in North America, 50-500 employees", + domains=["gusto.com", "rippling.com"], + integration_id=NATIVE_ICP_ENGINE, + ) +) +print(job.column_name) # ["ICP Fit", "ICP Score", "Reasoning"] +result = job.wait() +``` + +The engine decides the result columns, so read `job.column_name` instead of hardcoding them: an LLM run returns `Fit` / `Confidence` / `Reasoning`, the native model `ICP Fit` / `ICP Score` / `Reasoning (always null)`. `ICP Fit` is `Yes` or `No` at a 0.50 threshold on `ICP Score`, the calibrated probability as a 0.00-1.00 string. The native model returns only that score, so `Reasoning` is always `null` — the column is there to keep the set the same shape as an LLM run, not to carry an explanation. `integration_id="native-icp"` also works on `discogen.process` for a prompt that already carries the validation structure. + +Two errors are specific to the native engine: a 400 `ValidationError` when the ICP text does not yield a Mandatory / Reject if / Nice-to-have prompt, and a 503 `ServerError` when no ICP-fit engine is available. Task lifecycle, polling and statuses are the same either way. + ## Links - **API documentation**: [docs.discolike.com](https://docs.discolike.com) diff --git a/packages/discolike/pyproject.toml b/packages/discolike/pyproject.toml index f18d1e9..80784f7 100644 --- a/packages/discolike/pyproject.toml +++ b/packages/discolike/pyproject.toml @@ -34,7 +34,7 @@ classifiers = [ ] [project.optional-dependencies] -cli = ["discolike-cli==0.3.2"] +cli = ["discolike-cli==0.4.0"] [project.urls] Homepage = "https://www.discolike.com" diff --git a/packages/discolike/src/discolike/__init__.py b/packages/discolike/src/discolike/__init__.py index 6b3cbc7..f992306 100644 --- a/packages/discolike/src/discolike/__init__.py +++ b/packages/discolike/src/discolike/__init__.py @@ -18,6 +18,8 @@ from discolike._models import DiscolikeModel from discolike._models import DiscolikeRequest from discolike._version import __version__ +from discolike.resources.contacts import NATIVE_ENGINE +from discolike.resources.discogen import NATIVE_ICP_ENGINE from discolike.resources.discovery import Company from discolike.resources.discovery import Count from discolike.resources.email import EmailBatchResults @@ -30,6 +32,8 @@ from discolike.signup import signup __all__ = [ + "NATIVE_ENGINE", + "NATIVE_ICP_ENGINE", "APIConnectionError", "ApiKeyCredential", "AsyncDiscolike", diff --git a/packages/discolike/src/discolike/_credentials.py b/packages/discolike/src/discolike/_credentials.py index c08644a..80e9521 100644 --- a/packages/discolike/src/discolike/_credentials.py +++ b/packages/discolike/src/discolike/_credentials.py @@ -18,6 +18,9 @@ class OAuthCredential: expires_at: float client_id: str token_endpoint: str + # RFC 8707 resource the token was issued for; resent on refresh so an authorization server with a + # default resource cannot re-bind the refreshed token. None for credentials stored before 0.3.3. + resource: str | None = None def expires_within(self, seconds: float, *, now: float | None = None) -> bool: current = time.time() if now is None else now @@ -31,6 +34,7 @@ def from_config(cls, data: dict[str, Any]) -> OAuthCredential: expires_at=float(data["expires_at"]), client_id=str(data["client_id"]), token_endpoint=str(data["token_endpoint"]), + resource=str(data["resource"]) if data.get("resource") else None, ) def to_config(self) -> dict[str, Any]: diff --git a/packages/discolike/src/discolike/_generated/requests.py b/packages/discolike/src/discolike/_generated/requests.py index 59eb002..9f3b048 100644 --- a/packages/discolike/src/discolike/_generated/requests.py +++ b/packages/discolike/src/discolike/_generated/requests.py @@ -205,12 +205,15 @@ class ContactsSearchParams(DiscolikeRequest): ] = None filter_state: Annotated[ list[str] | None, - Field(description="Filter by company state/region.", title="Filter State"), + Field( + description="Filter by company state/region. Accepts ISO 3166-2 codes or names, listed at https://docs.discolike.com/states/, resolved against the selected countries; a value that resolves under none is ignored.", + title="Filter State", + ), ] = None negate_filter_state: Annotated[ list[str] | None, Field( - description="Exclude contacts at companies in specified states.", + description="Exclude contacts at companies in specified states. Same format as filter_state.", title="Negate Filter State", ), ] = None @@ -345,7 +348,10 @@ class ContactsSearchParams(DiscolikeRequest): ] = None person_state: Annotated[ list[str] | None, - Field(description="Filter by contact's state/region.", title="Person State"), + Field( + description="Filter by contact's state/region. Accepts ISO 3166-2 codes or names, listed at https://docs.discolike.com/states/, resolved against the selected countries; a value that resolves under none is ignored.", + title="Person State", + ), ] = None has_email: Annotated[ bool | None, @@ -628,12 +634,15 @@ class ContactsCountParams(DiscolikeRequest): ] = None filter_state: Annotated[ list[str] | None, - Field(description="Filter by company state/region.", title="Filter State"), + Field( + description="Filter by company state/region. Accepts ISO 3166-2 codes or names, listed at https://docs.discolike.com/states/, resolved against the selected countries; a value that resolves under none is ignored.", + title="Filter State", + ), ] = None negate_filter_state: Annotated[ list[str] | None, Field( - description="Exclude contacts at companies in specified states.", + description="Exclude contacts at companies in specified states. Same format as filter_state.", title="Negate Filter State", ), ] = None @@ -768,7 +777,10 @@ class ContactsCountParams(DiscolikeRequest): ] = None person_state: Annotated[ list[str] | None, - Field(description="Filter by contact's state/region.", title="Person State"), + Field( + description="Filter by contact's state/region. Accepts ISO 3166-2 codes or names, listed at https://docs.discolike.com/states/, resolved against the selected countries; a value that resolves under none is ignored.", + title="Person State", + ), ] = None has_email: Annotated[ bool | None, @@ -1083,12 +1095,15 @@ class ContactFilters(DiscolikeRequest): ] = None filter_state: Annotated[ list[str] | None, - Field(description="Filter by company state/region.", title="Filter State"), + Field( + description="Filter by company state/region. Accepts ISO 3166-2 codes or names, listed at https://docs.discolike.com/states/, resolved against the selected countries; a value that resolves under none is ignored.", + title="Filter State", + ), ] = None negate_filter_state: Annotated[ list[str] | None, Field( - description="Exclude contacts at companies in specified states.", + description="Exclude contacts at companies in specified states. Same format as filter_state.", title="Negate Filter State", ), ] = None @@ -1223,7 +1238,10 @@ class ContactFilters(DiscolikeRequest): ] = None person_state: Annotated[ list[str] | None, - Field(description="Filter by contact's state/region.", title="Person State"), + Field( + description="Filter by contact's state/region. Accepts ISO 3166-2 codes or names, listed at https://docs.discolike.com/states/, resolved against the selected countries; a value that resolves under none is ignored.", + title="Person State", + ), ] = None has_email: Annotated[ bool | None, @@ -1379,7 +1397,13 @@ class ContactGenerateRequest(DiscolikeRequest): ), ] context_mode: Annotated[Literal["website", "profile", "domain"] | None, Field(title="Context Mode")] = "website" - integration_id: Annotated[str | None, Field(title="Integration Id")] = None + integration_id: Annotated[ + str | None, + Field( + description="LLM provider integration UUID, or 'native' for keyless extraction without title validation. Omit for the org default (falls back to native when none is set).", + title="Integration Id", + ), + ] = None search_provider_id: Annotated[str | None, Field(title="Search Provider Id")] = None search_context_size: Annotated[Literal["low", "medium", "high"] | None, Field(title="Search Context Size")] = "low" max_contacts_per_domain: Annotated[int | None, Field(title="Max Contacts Per Domain")] = 10 @@ -1387,6 +1411,13 @@ class ContactGenerateRequest(DiscolikeRequest): full_domains: Annotated[list[str] | None, Field(title="Full Domains")] = None partial_domains: Annotated[list[str] | None, Field(title="Partial Domains")] = None initial_contact_counts: Annotated[dict[str, int] | None, Field(title="Initial Contact Counts")] = None + find_emails: Annotated[ + bool | None, + Field( + description="After discovery, run every contact that has a name but no email through the email finder (same as the app's Verify step) and fill `email` / `email_status` on the row before the task completes. Found addresses are billed under the finder's own rules; nothing else is.", + title="Find Emails", + ), + ] = False class DiscoGenProcessRequest(DiscolikeRequest): @@ -1397,12 +1428,20 @@ class DiscoGenProcessRequest(DiscolikeRequest): integration_id: Annotated[ str | None, Field( - description="LLM provider integration UUID (omit for org default)", + description="LLM provider integration UUID (omit for org default), or 'native-icp' to score an ICP validation prompt with the native model", title="Integration Id", ), ] = None web_search: Annotated[bool | None, Field(title="Web Search")] = False include_x_search: Annotated[bool | None, Field(title="Include X Search")] = False + typed_columns: Annotated[ + bool | None, + Field( + description="Let the detector answer yes/no, fixed-set and scale columns with a TypeSafe judgment model", + title="Typed Columns", + ), + ] = False + include_confidence: Annotated[bool | None, Field(title="Include Confidence")] = False search_provider_id: Annotated[ str | None, Field( @@ -1432,12 +1471,20 @@ class DiscoGenPersonaProcessRequest(DiscolikeRequest): integration_id: Annotated[ str | None, Field( - description="LLM provider integration UUID (omit for org default)", + description="LLM provider integration UUID (omit for org default), or 'native-icp' to score an ICP validation prompt with the native model", title="Integration Id", ), ] = None web_search: Annotated[bool | None, Field(title="Web Search")] = False include_x_search: Annotated[bool | None, Field(title="Include X Search")] = False + typed_columns: Annotated[ + bool | None, + Field( + description="Let the detector answer yes/no, fixed-set and scale columns with a TypeSafe judgment model", + title="Typed Columns", + ), + ] = False + include_confidence: Annotated[bool | None, Field(title="Include Confidence")] = False search_provider_id: Annotated[ str | None, Field( @@ -1483,7 +1530,7 @@ class ValidateIcpRequest(DiscolikeRequest): integration_id: Annotated[ str | None, Field( - description="LLM provider integration UUID (omit for org default)", + description="LLM provider integration UUID (omit for org default), or 'native-icp' to score with DiscoLike's own ICP-fit model at no LLM cost", title="Integration Id", ), ] = None @@ -1674,6 +1721,61 @@ class DiscoverParams(DiscolikeRequest): title="Negate Category", ), ] = None + sub_industry: Annotated[ + list[str] | None, + Field( + description="Filter by sub-industry, a second-level label scoped to an industry category. Accepts a bare label (ROOFING) or a parent-qualified key (CONSTRUCTION/ROOFING), case-insensitive, up to 50 values. A bare label whose parent category is unambiguous adds that parent to the category filter. Call list-industry-categories for the label list.", + max_length=50, + title="Sub Industry", + ), + ] = None + negate_sub_industry: Annotated[ + list[str] | None, + Field( + description="Exclude specified sub-industries. Same format as sub_industry; does not affect the category filter.", + max_length=50, + title="Negate Sub Industry", + ), + ] = None + lat: Annotated[ + float | None, + Field( + description="Latitude of the search centre. Must be supplied together with lon.", + ge=-90.0, + le=90.0, + title="Lat", + ), + ] = None + lon: Annotated[ + float | None, + Field( + description="Longitude of the search centre. Must be supplied together with lat.", + ge=-180.0, + le=180.0, + title="Lon", + ), + ] = None + radius: Annotated[ + str | None, + Field( + description="Search radius around lat/lon: a number optionally suffixed with km or mi (50km, 30mi, 50). A bare number is kilometres. Defaults to 50km when lat/lon are supplied, maximum 1000km.", + title="Radius", + ), + ] = None + geo: Annotated[ + list[str] | None, + Field( + description="A circular area to search, written lat,lon or lat,lon,radius (30.27,-97.74 or 30.27,-97.74,30mi). The radius is a number optionally suffixed with km or mi, a bare number meaning kilometres; it defaults to 50km and may not exceed 1000km. Repeatable: every geo circle, every bbox and the lat/lon/radius centre are OR'd together, up to 10 shapes in total.", + title="Geo", + ), + ] = None + bbox: Annotated[ + list[str] | None, + Field( + description="Bounding box as min_lat,min_lon,max_lat,max_lon. Longitudes may wrap the antimeridian (min_lon above max_lon). Repeatable: every geo circle, every bbox and the lat/lon/radius centre are OR'd together, up to 10 shapes in total.", + title="Bbox", + ), + ] = None min_digital_footprint: Annotated[ int | None, Field( @@ -1695,7 +1797,7 @@ class DiscoverParams(DiscolikeRequest): state: Annotated[ list[str] | None, Field( - description="Filter by state codes (up to 100). Not supported with multiple countries.", + description="Filter by ISO 3166-2 state codes or names (up to 100), listed at https://docs.discolike.com/states/. Requires exactly one country value, which may be a region alias; the state is then resolved against every country the alias covers. A value that resolves under none is rejected.", max_length=100, title="State", ), @@ -1703,7 +1805,7 @@ class DiscoverParams(DiscolikeRequest): negate_state: Annotated[ list[str] | None, Field( - description="Exclude specified states from results (up to 100).", + description="Exclude specified states from results (up to 100). Same format and country rules as state.", max_length=100, title="Negate State", ), @@ -2256,6 +2358,61 @@ class CountParams(DiscolikeRequest): title="Negate Category", ), ] = None + sub_industry: Annotated[ + list[str] | None, + Field( + description="Filter by sub-industry, a second-level label scoped to an industry category. Accepts a bare label (ROOFING) or a parent-qualified key (CONSTRUCTION/ROOFING), case-insensitive, up to 50 values. A bare label whose parent category is unambiguous adds that parent to the category filter. Call list-industry-categories for the label list.", + max_length=50, + title="Sub Industry", + ), + ] = None + negate_sub_industry: Annotated[ + list[str] | None, + Field( + description="Exclude specified sub-industries. Same format as sub_industry; does not affect the category filter.", + max_length=50, + title="Negate Sub Industry", + ), + ] = None + lat: Annotated[ + float | None, + Field( + description="Latitude of the search centre. Must be supplied together with lon.", + ge=-90.0, + le=90.0, + title="Lat", + ), + ] = None + lon: Annotated[ + float | None, + Field( + description="Longitude of the search centre. Must be supplied together with lat.", + ge=-180.0, + le=180.0, + title="Lon", + ), + ] = None + radius: Annotated[ + str | None, + Field( + description="Search radius around lat/lon: a number optionally suffixed with km or mi (50km, 30mi, 50). A bare number is kilometres. Defaults to 50km when lat/lon are supplied, maximum 1000km.", + title="Radius", + ), + ] = None + geo: Annotated[ + list[str] | None, + Field( + description="A circular area to search, written lat,lon or lat,lon,radius (30.27,-97.74 or 30.27,-97.74,30mi). The radius is a number optionally suffixed with km or mi, a bare number meaning kilometres; it defaults to 50km and may not exceed 1000km. Repeatable: every geo circle, every bbox and the lat/lon/radius centre are OR'd together, up to 10 shapes in total.", + title="Geo", + ), + ] = None + bbox: Annotated[ + list[str] | None, + Field( + description="Bounding box as min_lat,min_lon,max_lat,max_lon. Longitudes may wrap the antimeridian (min_lon above max_lon). Repeatable: every geo circle, every bbox and the lat/lon/radius centre are OR'd together, up to 10 shapes in total.", + title="Bbox", + ), + ] = None min_digital_footprint: Annotated[ int | None, Field( @@ -2277,7 +2434,7 @@ class CountParams(DiscolikeRequest): state: Annotated[ list[str] | None, Field( - description="Filter by state codes (up to 100). Not supported with multiple countries.", + description="Filter by ISO 3166-2 state codes or names (up to 100), listed at https://docs.discolike.com/states/. Requires exactly one country value, which may be a region alias; the state is then resolved against every country the alias covers. A value that resolves under none is rejected.", max_length=100, title="State", ), @@ -2285,7 +2442,7 @@ class CountParams(DiscolikeRequest): negate_state: Annotated[ list[str] | None, Field( - description="Exclude specified states from results (up to 100).", + description="Exclude specified states from results (up to 100). Same format and country rules as state.", max_length=100, title="Negate State", ), @@ -2552,6 +2709,27 @@ class FindEmailRequest(DiscolikeRequest): class FindEmailBatchRequest(DiscolikeRequest): + source_query_id: Annotated[ + str | None, + Field( + description="Saved query id whose stored contacts receive these verdicts", + title="Source Query Id", + ), + ] = None + refs: Annotated[ + list[str] | None, + Field( + description="Contact row ids, one per request, aligned by index (requires source_query_id)", + title="Refs", + ), + ] = None + round: Annotated[ + Literal["verify", "escalate"] | None, + Field( + description="verify = first pass, escalate = re-check of unproven", + title="Round", + ), + ] = "verify" requests: Annotated[ list[FindEmailRequest], Field( @@ -2650,7 +2828,10 @@ class MatchCompanyParams(DiscolikeRequest): ), ] = None city: Annotated[str | None, Field(description="City to augment the search", title="City")] = None - state: Annotated[str | None, Field(description="State code to augment the search", title="State")] = None + state: Annotated[ + str | None, + Field(description="State code or name to augment the search", title="State"), + ] = None country: Annotated[ str | None, Field( @@ -2699,7 +2880,10 @@ class MatchBulkParams(DiscolikeRequest): ] = None state_column: Annotated[ str | None, - Field(description="Column name containing states.", title="State Column"), + Field( + description="Column name containing states, as ISO 3166-2 codes or names.", + title="State Column", + ), ] = None country_column: Annotated[ str | None, diff --git a/packages/discolike/src/discolike/_jobs.py b/packages/discolike/src/discolike/_jobs.py index 83d3c35..b3f9355 100644 --- a/packages/discolike/src/discolike/_jobs.py +++ b/packages/discolike/src/discolike/_jobs.py @@ -4,6 +4,7 @@ import time from collections.abc import Callable from typing import Any +from typing import Literal import pydantic @@ -21,6 +22,11 @@ DEFAULT_WAIT_TIMEOUT_SECONDS = 900.0 DEFAULT_POLL_INTERVAL_SECONDS = 5.0 +# The engine that ran decides the result columns: an LLM validation returns Fit / Confidence / +# Reasoning, the native ICP-fit model ICP Fit / ICP Score / Reasoning, whose Reasoning is always +# null. Read it, never assume. +ColumnName = str | list[str] | None + class JobStatus(DiscolikeModel): status: str @@ -36,13 +42,15 @@ class JobStatus(DiscolikeModel): # model's built-in search only; on a BYOS run read search_provider instead. estimated_cost: float | None = None cost_metadata: dict[str, dict[str, Any]] | None = None + title_validation: Literal["llm", "none"] | None = None class Job: - def __init__(self, transport: Transport, *, task_family: str, task_id: str) -> None: + def __init__(self, transport: Transport, *, task_family: str, task_id: str, column_name: ColumnName = None) -> None: self._transport = transport self.task_family = task_family self.task_id = task_id + self.column_name = column_name def status(self) -> JobStatus: response = self._transport.request("GET", f"/{self.task_family}/status/{self.task_id}") @@ -76,10 +84,13 @@ def wait( class AsyncJob: - def __init__(self, transport: AsyncTransport, *, task_family: str, task_id: str) -> None: + def __init__( + self, transport: AsyncTransport, *, task_family: str, task_id: str, column_name: ColumnName = None + ) -> None: self._transport = transport self.task_family = task_family self.task_id = task_id + self.column_name = column_name async def status(self) -> JobStatus: response = await self._transport.request("GET", f"/{self.task_family}/status/{self.task_id}") diff --git a/packages/discolike/src/discolike/_models.py b/packages/discolike/src/discolike/_models.py index d400be6..914f48a 100644 --- a/packages/discolike/src/discolike/_models.py +++ b/packages/discolike/src/discolike/_models.py @@ -1,9 +1,28 @@ from __future__ import annotations +import math +from collections.abc import Callable from typing import Any import pydantic +BBOX_FIELD_COUNT = 4 +BBOX_SEPARATOR = "," +BBOX_FORMAT = "min_lat,min_lon,max_lat,max_lon" +GEO_SEPARATOR = "," +GEO_FORMAT = "lat,lon or lat,lon,radius" +GEO_POINT_FIELD_COUNT = 2 +GEO_CIRCLE_FIELD_COUNT = 3 +MIN_LATITUDE = -90.0 +MAX_LATITUDE = 90.0 +MIN_LONGITUDE = -180.0 +MAX_LONGITUDE = 180.0 +MAX_RADIUS_KM = 1000.0 +KM_PER_MILE = 1.609344 +RADIUS_KM_SUFFIX = "km" +RADIUS_MILE_SUFFIX = "mi" +MAX_GEO_SHAPES = 10 + class DiscolikeModel(pydantic.BaseModel): model_config = pydantic.ConfigDict(extra="allow") @@ -15,5 +34,110 @@ def to_dict(self) -> dict[str, Any]: class DiscolikeRequest(pydantic.BaseModel): model_config = pydantic.ConfigDict(extra="allow", populate_by_name=True) + # Mirrors the platform's parse_bbox and parse_geo_circle so a bad shape fails here instead of as + # an opaque 4xx. The request models are generated, so the validators can only be attached from + # the base class; check_fields=False keeps them inert on the models that have no geo or bbox. + # A lone string is wrapped into a one-element list the way the platform coerces it. + @pydantic.field_validator("bbox", mode="before", check_fields=False) + @classmethod + def _validate_bbox(cls, value: object) -> object: + return _validate_shapes(value=value, name="bbox", fmt=BBOX_FORMAT, rejection=_bbox_rejection) + + @pydantic.field_validator("geo", mode="before", check_fields=False) + @classmethod + def _validate_geo(cls, value: object) -> object: + return _validate_shapes(value=value, name="geo", fmt=GEO_FORMAT, rejection=_geo_rejection) + + @pydantic.model_validator(mode="after") + def _validate_geo_shapes(self) -> DiscolikeRequest: + if "lat" not in type(self).model_fields: + return self + lat, lon, radius = getattr(self, "lat", None), getattr(self, "lon", None), getattr(self, "radius", None) + if (lat is None) != (lon is None): + raise ValueError("lat and lon must be supplied together") + if radius is not None: + if lat is None: + raise ValueError("radius needs lat and lon") + reason = _radius_km_rejection(radius) + if reason is not None: + raise ValueError(reason) + total = (lat is not None) + len(getattr(self, "geo", None) or []) + len(getattr(self, "bbox", None) or []) + if total > MAX_GEO_SHAPES: + raise ValueError(f"{total} geo shapes (lat/lon, geo and bbox together); at most {MAX_GEO_SHAPES}") + return self + def to_wire(self) -> dict[str, Any]: return self.model_dump(mode="json", exclude_unset=True, by_alias=True) + + +def _bbox_rejection(value: str) -> str | None: + parts = value.replace(" ", "").split(BBOX_SEPARATOR) + if len(parts) != BBOX_FIELD_COUNT: + return f"expected {BBOX_FIELD_COUNT} values, got {len(parts)}" + try: + min_lat, min_lon, max_lat, max_lon = (float(part) for part in parts) + except ValueError: + return "every value must be a number" + if not all(math.isfinite(corner) for corner in (min_lat, min_lon, max_lat, max_lon)): + return "every value must be finite" + if not MIN_LATITUDE <= min_lat < max_lat <= MAX_LATITUDE: + return f"latitudes must satisfy {MIN_LATITUDE} <= min_lat < max_lat <= {MAX_LATITUDE}" + if not (MIN_LONGITUDE <= min_lon <= MAX_LONGITUDE and MIN_LONGITUDE <= max_lon <= MAX_LONGITUDE): + return f"longitudes must be between {MIN_LONGITUDE} and {MAX_LONGITUDE}" + if min_lon == max_lon: + return "min_lon and max_lon must differ" + return None + + +def _validate_shapes( + *, + value: object, + name: str, + fmt: str, + rejection: Callable[[str], str | None], +) -> object: + if value is None or not isinstance(value, str | list | tuple): + return value + shapes = [value] if isinstance(value, str) else list(value) + for shape in shapes: + if not isinstance(shape, str): + return value + reason = rejection(shape) + if reason is not None: + raise ValueError(f"invalid {name} {shape!r}: {reason}; use {fmt}") + return shapes + + +def _radius_km_rejection(value: str) -> str | None: + raw = value.strip().lower().replace(" ", "") + multiplier = 1.0 + if raw.endswith(RADIUS_MILE_SUFFIX): + raw, multiplier = raw[: -len(RADIUS_MILE_SUFFIX)], KM_PER_MILE + elif raw.endswith(RADIUS_KM_SUFFIX): + raw = raw[: -len(RADIUS_KM_SUFFIX)] + try: + radius_km = float(raw) * multiplier + except ValueError: + return f"radius {value!r} must be a number optionally suffixed with {RADIUS_KM_SUFFIX} or {RADIUS_MILE_SUFFIX}" + if not math.isfinite(radius_km) or radius_km <= 0 or radius_km > MAX_RADIUS_KM: + return f"radius {value!r} must be greater than 0 and at most {MAX_RADIUS_KM:g}km" + return None + + +def _geo_rejection(value: str) -> str | None: + parts = value.replace(" ", "").split(GEO_SEPARATOR) + if not GEO_POINT_FIELD_COUNT <= len(parts) <= GEO_CIRCLE_FIELD_COUNT: + return f"expected {GEO_POINT_FIELD_COUNT} or {GEO_CIRCLE_FIELD_COUNT} values, got {len(parts)}" + try: + lat, lon = float(parts[0]), float(parts[1]) + except ValueError: + return "latitude and longitude must be numbers" + if not (math.isfinite(lat) and math.isfinite(lon)): + return "latitude and longitude must be finite" + if not MIN_LATITUDE <= lat <= MAX_LATITUDE: + return f"latitude must be between {MIN_LATITUDE} and {MAX_LATITUDE}" + if not MIN_LONGITUDE <= lon <= MAX_LONGITUDE: + return f"longitude must be between {MIN_LONGITUDE} and {MAX_LONGITUDE}" + if len(parts) == GEO_CIRCLE_FIELD_COUNT: + return _radius_km_rejection(parts[-1]) + return None diff --git a/packages/discolike/src/discolike/_oauth.py b/packages/discolike/src/discolike/_oauth.py index 7289cad..9e39663 100644 --- a/packages/discolike/src/discolike/_oauth.py +++ b/packages/discolike/src/discolike/_oauth.py @@ -111,7 +111,12 @@ def _client_kwargs( def _credential_from_token( - token: dict[str, Any], *, client_id: str, token_endpoint: str, fallback_refresh_token: str | None + token: dict[str, Any], + *, + client_id: str, + token_endpoint: str, + resource: str | None, + fallback_refresh_token: str | None, ) -> OAuthCredential: # Exceptions raised here may be logged by SDK consumers; never attach live tokens to them. safe_payload = {key: value for key, value in token.items() if key not in TOKEN_KEYS} @@ -128,9 +133,18 @@ def _credential_from_token( expires_at=float(token["expires_at"]), client_id=client_id, token_endpoint=token_endpoint, + resource=resource, ) +def _refresh_kwargs(credential: OAuthCredential) -> dict[str, Any]: + """Resend the resource the token was issued for; credentials stored before 0.3.3 have none.""" + kwargs: dict[str, Any] = {"refresh_token": credential.refresh_token} + if credential.resource: + kwargs["resource"] = credential.resource + return kwargs + + def discover(base_url: str, *, client: httpx2.Client) -> AuthServerMetadata: payload = _payload(client.get(base_url.rstrip("/") + METADATA_PATH)) return AuthServerMetadata( @@ -178,7 +192,11 @@ def exchange_code( with TokenClient(**kwargs) as client: token = client.fetch_token(metadata.token_endpoint, code=code, code_verifier=code_verifier, resource=resource) return _credential_from_token( - token, client_id=client_id, token_endpoint=metadata.token_endpoint, fallback_refresh_token=None + token, + client_id=client_id, + token_endpoint=metadata.token_endpoint, + resource=resource, + fallback_refresh_token=None, ) @@ -188,11 +206,12 @@ def refresh( """Rotates the tokens; any failure means the session is gone and the user must log in again.""" try: with TokenClient(**_client_kwargs(client_id=credential.client_id, transport=transport)) as client: - token = client.refresh_token(credential.token_endpoint, refresh_token=credential.refresh_token) + token = client.refresh_token(credential.token_endpoint, **_refresh_kwargs(credential)) return _credential_from_token( token, client_id=credential.client_id, token_endpoint=credential.token_endpoint, + resource=credential.resource, fallback_refresh_token=credential.refresh_token, ) except AuthenticationError as exc: @@ -204,11 +223,12 @@ async def refresh_async( ) -> OAuthCredential: try: async with AsyncTokenClient(**_client_kwargs(client_id=credential.client_id, transport=transport)) as client: - token = await client.refresh_token(credential.token_endpoint, refresh_token=credential.refresh_token) + token = await client.refresh_token(credential.token_endpoint, **_refresh_kwargs(credential)) return _credential_from_token( token, client_id=credential.client_id, token_endpoint=credential.token_endpoint, + resource=credential.resource, fallback_refresh_token=credential.refresh_token, ) except AuthenticationError as exc: diff --git a/packages/discolike/src/discolike/_version.py b/packages/discolike/src/discolike/_version.py index f9aa3e1..6a9beea 100644 --- a/packages/discolike/src/discolike/_version.py +++ b/packages/discolike/src/discolike/_version.py @@ -1 +1 @@ -__version__ = "0.3.2" +__version__ = "0.4.0" diff --git a/packages/discolike/src/discolike/resources/companies.py b/packages/discolike/src/discolike/resources/companies.py index 5d62bf0..4f1cc91 100644 --- a/packages/discolike/src/discolike/resources/companies.py +++ b/packages/discolike/src/discolike/resources/companies.py @@ -39,6 +39,10 @@ class CompanyProfile(DiscolikeModel): start_date: str | None = None end_date: str | None = None address: CompanyAddress | None = None + # The API returned these as lat/lon before the release that renamed them. + latitude: float | None = pydantic.Field(default=None, validation_alias=pydantic.AliasChoices("latitude", "lat")) + longitude: float | None = pydantic.Field(default=None, validation_alias=pydantic.AliasChoices("longitude", "lon")) + geo_precision: str | None = None phones: list[str] | None = None public_emails: list[str] | None = None domain_associations: list[str] = pydantic.Field(default_factory=list) @@ -47,6 +51,7 @@ class CompanyProfile(DiscolikeModel): description: str | None = None keywords: dict[str, float] = pydantic.Field(default_factory=dict) industry_groups: dict[str, float] = pydantic.Field(default_factory=dict) + sub_industry: dict[str, float] | None = None employees: str | None = None revenue_range: str | None = None business_model: dict[str, float] = pydantic.Field(default_factory=dict) diff --git a/packages/discolike/src/discolike/resources/contacts.py b/packages/discolike/src/discolike/resources/contacts.py index 98c801d..d284397 100644 --- a/packages/discolike/src/discolike/resources/contacts.py +++ b/packages/discolike/src/discolike/resources/contacts.py @@ -20,6 +20,8 @@ from discolike.resources.companies import CompanyProfile from discolike.resources.discovery import Count +NATIVE_ENGINE = "native" + class Contact(DiscolikeModel): persona_id: int | None = None diff --git a/packages/discolike/src/discolike/resources/discogen.py b/packages/discolike/src/discolike/resources/discogen.py index dae85f6..d6719bd 100644 --- a/packages/discolike/src/discolike/resources/discogen.py +++ b/packages/discolike/src/discolike/resources/discogen.py @@ -1,11 +1,14 @@ from __future__ import annotations +import httpx2 import pydantic from discolike._jobs import FAMILY_DISCOGEN from discolike._jobs import AsyncJob from discolike._jobs import Job from discolike._models import DiscolikeModel +from discolike._transport import AsyncTransport +from discolike._transport import Transport from discolike.requests import DiscoGenPersonaProcessRequest from discolike.requests import DiscoGenProcessRequest from discolike.requests import ValidateIcpRequest @@ -13,6 +16,28 @@ from discolike.resources._base import SyncAPIResource from discolike.resources._base import api_route +NATIVE_ICP_ENGINE = "native-icp" + + +def _job(transport: Transport, response: httpx2.Response) -> Job: + payload = response.json() + return Job( + transport, + task_family=FAMILY_DISCOGEN, + task_id=payload["task_id"], + column_name=payload.get("column_name"), + ) + + +def _async_job(transport: AsyncTransport, response: httpx2.Response) -> AsyncJob: + payload = response.json() + return AsyncJob( + transport, + task_family=FAMILY_DISCOGEN, + task_id=payload["task_id"], + column_name=payload.get("column_name"), + ) + class DiscogenModelInfo(DiscolikeModel): name: str | None = None @@ -26,13 +51,17 @@ class DiscogenModels(DiscolikeModel): class DiscogenResource(SyncAPIResource): @api_route("POST", "/discogen/process") def process(self, request: DiscoGenProcessRequest) -> Job: - response = self._transport.request("POST", "/discogen/process", json_body=request.to_wire()) - return Job(self._transport, task_family=FAMILY_DISCOGEN, task_id=response.json()["task_id"]) + """Run a research prompt against a list of domains. + + `integration_id` takes an LLM provider integration UUID, or `NATIVE_ICP_ENGINE` to score + an ICP validation prompt with DiscoLike's own model. Omit it for the org default. + """ + return _job(self._transport, self._transport.request("POST", "/discogen/process", json_body=request.to_wire())) @api_route("POST", "/discogen/process-personas") def process_personas(self, request: DiscoGenPersonaProcessRequest) -> Job: response = self._transport.request("POST", "/discogen/process-personas", json_body=request.to_wire()) - return Job(self._transport, task_family=FAMILY_DISCOGEN, task_id=response.json()["task_id"]) + return _job(self._transport, response) @api_route("GET", "/discogen/models") def models(self) -> DiscogenModels: @@ -45,20 +74,32 @@ def job(self, task_id: str) -> Job: class ValidateResource(SyncAPIResource): @api_route("POST", "/validate/icp") def icp(self, request: ValidateIcpRequest) -> Job: - response = self._transport.request("POST", "/validate/icp", json_body=request.to_wire()) - return Job(self._transport, task_family=FAMILY_DISCOGEN, task_id=response.json()["task_id"]) + """Validate domains against an ICP description. + + `integration_id` takes an LLM provider integration UUID, or `NATIVE_ICP_ENGINE` to score + with DiscoLike's own ICP-fit model: no LLM key, no LLM cost, and no web search on that + run. Omit it for the org default. Read `Job.column_name` for the columns the run returns + rather than assuming the LLM ones; the native run's `Reasoning` column is always null, + since the model returns a score rather than an explanation. + """ + return _job(self._transport, self._transport.request("POST", "/validate/icp", json_body=request.to_wire())) class AsyncDiscogenResource(AsyncAPIResource): @api_route("POST", "/discogen/process") async def process(self, request: DiscoGenProcessRequest) -> AsyncJob: + """Run a research prompt against a list of domains. + + `integration_id` takes an LLM provider integration UUID, or `NATIVE_ICP_ENGINE` to score + an ICP validation prompt with DiscoLike's own model. Omit it for the org default. + """ response = await self._transport.request("POST", "/discogen/process", json_body=request.to_wire()) - return AsyncJob(self._transport, task_family=FAMILY_DISCOGEN, task_id=response.json()["task_id"]) + return _async_job(self._transport, response) @api_route("POST", "/discogen/process-personas") async def process_personas(self, request: DiscoGenPersonaProcessRequest) -> AsyncJob: response = await self._transport.request("POST", "/discogen/process-personas", json_body=request.to_wire()) - return AsyncJob(self._transport, task_family=FAMILY_DISCOGEN, task_id=response.json()["task_id"]) + return _async_job(self._transport, response) @api_route("GET", "/discogen/models") async def models(self) -> DiscogenModels: @@ -72,5 +113,13 @@ def job(self, task_id: str) -> AsyncJob: class AsyncValidateResource(AsyncAPIResource): @api_route("POST", "/validate/icp") async def icp(self, request: ValidateIcpRequest) -> AsyncJob: + """Validate domains against an ICP description. + + `integration_id` takes an LLM provider integration UUID, or `NATIVE_ICP_ENGINE` to score + with DiscoLike's own ICP-fit model: no LLM key, no LLM cost, and no web search on that + run. Omit it for the org default. Read `AsyncJob.column_name` for the columns the run + returns rather than assuming the LLM ones; the native run's `Reasoning` column is always + null, since the model returns a score rather than an explanation. + """ response = await self._transport.request("POST", "/validate/icp", json_body=request.to_wire()) - return AsyncJob(self._transport, task_family=FAMILY_DISCOGEN, task_id=response.json()["task_id"]) + return _async_job(self._transport, response) diff --git a/packages/discolike/tests/test_config.py b/packages/discolike/tests/test_config.py index 5daa832..7f28b50 100644 --- a/packages/discolike/tests/test_config.py +++ b/packages/discolike/tests/test_config.py @@ -60,7 +60,12 @@ def test_binary_garbage_config_returns_empty(isolated_config) -> None: def _oauth_credential() -> OAuthCredential: return OAuthCredential( - access_token="at", refresh_token="rt", expires_at=1.0, client_id="c", token_endpoint="https://t/token" + access_token="at", + refresh_token="rt", + expires_at=1.0, + client_id="c", + token_endpoint="https://t/token", + resource="https://api.example.com/v1", ) @@ -74,10 +79,30 @@ def test_save_and_load_oauth_credential(isolated_config) -> None: "expires_at": 1.0, "client_id": "c", "token_endpoint": "https://t/token", + "resource": "https://api.example.com/v1", } assert load_credential() == _oauth_credential() +def test_load_oauth_credential_saved_before_resource_was_stored(isolated_config) -> None: + """Config written by 0.3.x has no `resource`; it must still load, and refresh then omits the field.""" + save_config( + { + "auth_method": "oauth", + "oauth": { + "access_token": "at", + "refresh_token": "rt", + "expires_at": 1.0, + "client_id": "c", + "token_endpoint": "https://t/token", + }, + } + ) + loaded = load_credential() + assert isinstance(loaded, OAuthCredential) + assert loaded.resource is None + + def test_save_api_key_credential_keeps_legacy_shape(isolated_config) -> None: save_credential(ApiKeyCredential(api_key="dk-1")) assert load_config() == {"auth_method": "api_key", "api_key": "dk-1"} diff --git a/packages/discolike/tests/test_contacts.py b/packages/discolike/tests/test_contacts.py index ba13d2c..7f8765e 100644 --- a/packages/discolike/tests/test_contacts.py +++ b/packages/discolike/tests/test_contacts.py @@ -17,6 +17,7 @@ from discolike.requests import ContactsLookupParams from discolike.requests import ContactsMatchParams from discolike.requests import ContactsSearchParams +from discolike.resources.contacts import NATIVE_ENGINE from discolike_testkit import AsyncClientFactory from discolike_testkit import ClientFactory @@ -237,6 +238,65 @@ def handler(request: httpx2.Request) -> httpx2.Response: assert job.task_id == "dg-1" +def test_generate_native_engine_serializes_integration_id(make_client: ClientFactory) -> None: + seen = {} + + def handler(request: httpx2.Request) -> httpx2.Response: + seen["body"] = json.loads(request.content) + return httpx2.Response(200, json={"task_id": "dg-native"}) + + with make_client(handler) as client: + job = client.contacts.generate( + ContactGenerateRequest( + icp_text="VPs of Marketing at B2B SaaS", + domains=["gusto.com"], + integration_id=NATIVE_ENGINE, + ) + ) + + assert NATIVE_ENGINE == "native" + assert seen["body"]["integration_id"] == "native" + assert job.task_id == "dg-native" + + +async def test_generate_async_native_engine_serializes_integration_id( + make_async_client: AsyncClientFactory, +) -> None: + seen = {} + + def handler(request: httpx2.Request) -> httpx2.Response: + seen["body"] = json.loads(request.content) + return httpx2.Response(200, json={"task_id": "dg-native-2"}) + + async with make_async_client(handler) as client: + job = await client.contacts.generate( + ContactGenerateRequest(icp_text="VPs", domains=["a.com"], integration_id=NATIVE_ENGINE) + ) + + assert seen["body"]["integration_id"] == "native" + assert job.task_id == "dg-native-2" + + +def test_generate_omits_integration_id_when_unset(make_client: ClientFactory) -> None: + seen = {} + + def handler(request: httpx2.Request) -> httpx2.Response: + seen["body"] = json.loads(request.content) + return httpx2.Response(200, json={"task_id": "dg-default"}) + + with make_client(handler) as client: + client.contacts.generate(ContactGenerateRequest(icp_text="VPs", domains=["a.com"])) + + assert "integration_id" not in seen["body"] + + +def test_native_engine_is_exported_from_package_root() -> None: + import discolike + + assert discolike.NATIVE_ENGINE == NATIVE_ENGINE + assert "NATIVE_ENGINE" in discolike.__all__ + + async def test_search_async(make_async_client: AsyncClientFactory) -> None: def handler(request: httpx2.Request) -> httpx2.Response: return httpx2.Response(200, json=[{"persona_id": 2, "domain": "b.com"}]) diff --git a/packages/discolike/tests/test_discogen.py b/packages/discolike/tests/test_discogen.py index 67e0bab..20423a7 100644 --- a/packages/discolike/tests/test_discogen.py +++ b/packages/discolike/tests/test_discogen.py @@ -12,6 +12,7 @@ from discolike.requests import DiscoGenPersonaProcessRequest from discolike.requests import DiscoGenProcessRequest from discolike.requests import ValidateIcpRequest +from discolike.resources.discogen import NATIVE_ICP_ENGINE from discolike_testkit import AsyncClientFactory from discolike_testkit import ClientFactory @@ -71,6 +72,8 @@ def handler(request: httpx2.Request) -> httpx2.Response: web_search=True, context_mode="website", include_x_search=False, + typed_columns=True, + include_confidence=True, search_provider_id="serper", search_context_size="medium", ) @@ -83,11 +86,26 @@ def handler(request: httpx2.Request) -> httpx2.Response: "web_search": True, "context_mode": "website", "include_x_search": False, + "typed_columns": True, + "include_confidence": True, "search_provider_id": "serper", "search_context_size": "medium", } +def test_process_sends_typed_columns_when_set(make_client: ClientFactory) -> None: + seen = {} + + def handler(request: httpx2.Request) -> httpx2.Response: + seen["body"] = json.loads(request.content) + return httpx2.Response(200, json={"task_id": "dg-typed"}) + + with make_client(handler) as client: + client.discogen.process(DiscoGenProcessRequest(query="q", domains=["a.com"], typed_columns=True)) + + assert seen["body"] == {"query": "q", "domains": ["a.com"], "typed_columns": True} + + def test_process_personas_posts_json_and_returns_job(make_client: ClientFactory) -> None: seen = {} @@ -280,3 +298,108 @@ def test_route_metadata_stamped() -> None: assert get_discolike_route(DiscogenResource.models) == ("GET", "/discogen/models", True) assert get_discolike_route(DiscogenResource.job) is None assert get_discolike_route(ValidateResource.icp) == ("POST", "/validate/icp", True) + + +def test_validate_icp_native_sentinel_serializes_and_exposes_columns(make_client: ClientFactory) -> None: + seen = {} + + def handler(request: httpx2.Request) -> httpx2.Response: + seen["body"] = json.loads(request.content) + return httpx2.Response( + 200, + json={ + "task_id": "val-native", + "column_name": ["ICP Fit", "ICP Score", "Reasoning"], + "status": "in_progress", + "total_domains": 1, + }, + ) + + with make_client(handler) as client: + job = client.validate_icp( + ValidateIcpRequest( + icp_text="Cybersecurity for SMBs", + domains=["gusto.com"], + integration_id=NATIVE_ICP_ENGINE, + ) + ) + + assert NATIVE_ICP_ENGINE == "native-icp" + assert seen["body"]["integration_id"] == "native-icp" + assert job.task_id == "val-native" + assert job.column_name == ["ICP Fit", "ICP Score", "Reasoning"] + + +async def test_validate_icp_async_native_sentinel_serializes_and_exposes_columns( + make_async_client: AsyncClientFactory, +) -> None: + seen = {} + + def handler(request: httpx2.Request) -> httpx2.Response: + seen["body"] = json.loads(request.content) + return httpx2.Response( + 200, json={"task_id": "val-native-async", "column_name": ["ICP Fit", "ICP Score", "Reasoning"]} + ) + + async with make_async_client(handler) as client: + job = await client.validate_icp( + ValidateIcpRequest(icp_text="q", domains=["a.com"], integration_id=NATIVE_ICP_ENGINE) + ) + + assert seen["body"]["integration_id"] == "native-icp" + assert job.column_name == ["ICP Fit", "ICP Score", "Reasoning"] + + +def test_validate_icp_llm_run_reports_its_own_columns(make_client: ClientFactory) -> None: + def handler(request: httpx2.Request) -> httpx2.Response: + return httpx2.Response(200, json={"task_id": "val-llm", "column_name": ["Fit", "Confidence", "Reasoning"]}) + + with make_client(handler) as client: + job = client.validate_icp(ValidateIcpRequest(icp_text="q", domains=["a.com"], integration_id="int-1")) + + assert job.column_name == ["Fit", "Confidence", "Reasoning"] + + +def test_validate_icp_column_name_is_none_when_absent(make_client: ClientFactory) -> None: + def handler(request: httpx2.Request) -> httpx2.Response: + return httpx2.Response(200, json={"task_id": "val-bare"}) + + with make_client(handler) as client: + job = client.validate_icp(ValidateIcpRequest(icp_text="q", domains=["a.com"])) + + assert job.column_name is None + + +def test_process_native_icp_sentinel_serializes(make_client: ClientFactory) -> None: + seen = {} + + def handler(request: httpx2.Request) -> httpx2.Response: + seen["body"] = json.loads(request.content) + return httpx2.Response(200, json={"task_id": "dg-native-icp", "column_name": ["ICP Fit"]}) + + with make_client(handler) as client: + job = client.discogen.process( + DiscoGenProcessRequest( + query="Mandatory:\nUS-based\n\nReject if:\nagency", + domains=["a.com"], + integration_id=NATIVE_ICP_ENGINE, + ) + ) + + assert seen["body"]["integration_id"] == "native-icp" + assert job.column_name == ["ICP Fit"] + + +def test_native_icp_engine_is_exported_from_package_root() -> None: + import discolike + + assert discolike.NATIVE_ICP_ENGINE == NATIVE_ICP_ENGINE + assert "NATIVE_ICP_ENGINE" in discolike.__all__ + + +def test_reattached_job_has_no_column_name(make_client: ClientFactory) -> None: + def handler(request: httpx2.Request) -> httpx2.Response: + pytest.fail("job() must not perform an HTTP request") + + with make_client(handler) as client: + assert client.discogen.job("dg-existing").column_name is None diff --git a/packages/discolike/tests/test_gen_requests.py b/packages/discolike/tests/test_gen_requests.py index 26d5421..93f4630 100644 --- a/packages/discolike/tests/test_gen_requests.py +++ b/packages/discolike/tests/test_gen_requests.py @@ -214,6 +214,29 @@ def test_normalize_schema_strips_scalar_item_constraints(gen) -> None: } +def test_apply_overlays_fills_only_what_the_spec_is_missing(gen) -> None: + kept = {"DiscoverParams": {"type": "object", "properties": {"lat": {"type": "number", "description": "deployed"}}}} + + properties = gen.apply_overlays(kept=kept)["DiscoverParams"]["properties"] + + assert properties["lat"] == {"type": "number", "description": "deployed"} + assert properties["radius"]["type"] == "string" + assert properties["sub_industry"]["items"] == {"type": "string"} + + +def test_apply_overlays_pins_sub_industry_over_a_spec_enum(gen) -> None: + kept = {"CountParams": {"type": "object", "properties": {"sub_industry": {"enum": ["CONSTRUCTION/ROOFING"]}}}} + + sub_industry = gen.apply_overlays(kept=kept)["CountParams"]["properties"]["sub_industry"] + + assert "enum" not in sub_industry + assert sub_industry["items"] == {"type": "string"} + + +def test_apply_overlays_skips_schemas_this_run_does_not_generate(gen) -> None: + assert gen.apply_overlays(kept={}) == {} + + def test_build_codegen_spec_wraps_pruned_schemas(gen, routes) -> None: codegen_spec = gen.build_codegen_spec(spec=FAKE_SPEC, routes=routes) assert codegen_spec["paths"] == {} diff --git a/packages/discolike/tests/test_jobs.py b/packages/discolike/tests/test_jobs.py index d71a463..cba2095 100644 --- a/packages/discolike/tests/test_jobs.py +++ b/packages/discolike/tests/test_jobs.py @@ -59,7 +59,7 @@ def test_status_exposes_cost_metadata_and_warnings() -> None: payload = { "status": "completed", "progress": 100, - "results": {"a.com": "yes"}, + "results": {"a.com": "Yes"}, "estimated_cost": 0.0283, "warnings": ["Search provider out of credits"], "cost_metadata": { @@ -88,6 +88,19 @@ def test_status_without_cost_fields_defaults_to_none() -> None: assert final.warnings == [] +def test_status_exposes_title_validation() -> None: + native = make_job(_status_sequence([{"status": "completed", "progress": 100, "title_validation": "none"}])).status() + assert native.title_validation == "none" + + byok = make_job(_status_sequence([{"status": "completed", "progress": 100, "title_validation": "llm"}])).status() + assert byok.title_validation == "llm" + + +def test_status_without_title_validation_defaults_to_none() -> None: + final = make_job(_status_sequence([{"status": "in_progress", "progress": 10}])).status() + assert final.title_validation is None + + def test_wait_failed_raises() -> None: handler = _status_sequence([{"status": "failed", "progress": 100, "result": "LLM exploded"}]) with pytest.raises(JobFailedError, match="LLM exploded"): diff --git a/packages/discolike/tests/test_models.py b/packages/discolike/tests/test_models.py index 1b8b12a..044a554 100644 --- a/packages/discolike/tests/test_models.py +++ b/packages/discolike/tests/test_models.py @@ -2,6 +2,8 @@ import pytest from discolike._models import DiscolikeRequest +from discolike.requests import CountParams +from discolike.requests import DiscoverParams class _Probe(DiscolikeRequest): @@ -36,3 +38,135 @@ def test_discolike_request_is_exported_from_the_package() -> None: import discolike assert discolike.DiscolikeRequest is DiscolikeRequest + + +class _BboxProbe(DiscolikeRequest): + bbox: list[str] | None = None + + +class _GeoProbe(DiscolikeRequest): + geo: list[str] | None = None + + +@pytest.mark.parametrize( + "value", + [ + "40.4,-74.3,41.0,-73.7", + "-90,-180,90,180", + "-10,170,10,-170", + "40.4, -74.3, 41.0, -73.7", + ], +) +def test_bbox_accepts_valid_boxes(value: str) -> None: + assert _BboxProbe(bbox=value).to_wire() == {"bbox": [value]} + + +def test_bbox_accepts_several_boxes() -> None: + boxes = ["40.4,-74.3,41.0,-73.7", "51.2,-0.5,51.7,0.3"] + assert _BboxProbe(bbox=boxes).to_wire() == {"bbox": boxes} + + +def test_bbox_rejects_a_bad_box_among_good_ones() -> None: + with pytest.raises(pydantic.ValidationError, match="latitudes must satisfy"): + _BboxProbe(bbox=["40.4,-74.3,41.0,-73.7", "-91,0,10,1"]) + + +@pytest.mark.parametrize( + "value", + [ + "30.27,-97.74", + "30.27,-97.74,30mi", + "30.27,-97.74,10km", + "30.27,-97.74,10", + "30.27, -97.74, 10 km", + "-90,-180", + "90,180,1000km", + ], +) +def test_geo_accepts_valid_circles(value: str) -> None: + assert _GeoProbe(geo=value).to_wire() == {"geo": [value]} + + +def test_geo_accepts_several_circles() -> None: + circles = ["30.27,-97.74,10km", "52.52,13.405"] + assert _GeoProbe(geo=circles).to_wire() == {"geo": circles} + + +@pytest.mark.parametrize( + ("value", "message"), + [ + ("30.27", "expected 2 or 3 values"), + ("30.27,-97.74,10km,1", "expected 2 or 3 values"), + ("north,-97.74", "must be numbers"), + ("nan,-97.74", "must be finite"), + ("90.5,-97.74", "latitude must be between"), + ("30.27,-180.5", "longitude must be between"), + ("30.27,-97.74,0km", "greater than 0"), + ("30.27,-97.74,1001km", "greater than 0"), + ("30.27,-97.74,wide", "must be a number optionally suffixed"), + ], +) +def test_geo_rejects_circles_the_platform_would_reject(value: str, message: str) -> None: + with pytest.raises(pydantic.ValidationError, match=message): + _GeoProbe(geo=value) + + +def test_geo_none_stays_unvalidated() -> None: + assert _GeoProbe(geo=None).to_wire() == {"geo": None} + + +@pytest.mark.parametrize( + ("value", "message"), + [ + ("40.4,-74.3,41.0", "expected 4 values"), + ("40.4,-74.3,41.0,-73.7,1", "expected 4 values"), + ("40.4,-74.3,41.0,east", "must be a number"), + ("nan,-74.3,41.0,-73.7", "must be finite"), + ("-inf,-74.3,41.0,-73.7", "must be finite"), + ("-91,0,10,1", "latitudes must satisfy"), + ("0,0,90.5,1", "latitudes must satisfy"), + ("41,-74,40,-73", "latitudes must satisfy"), + ("0,0,0,1", "latitudes must satisfy"), + ("0,-180.5,10,1", "longitudes must be between"), + ("0,0,10,180.5", "longitudes must be between"), + ("0,10,1,10", "min_lon and max_lon must differ"), + ], +) +def test_bbox_rejects_boxes_the_platform_would_reject(value: str, message: str) -> None: + with pytest.raises(pydantic.ValidationError, match=message): + _BboxProbe(bbox=value) + + +def test_bbox_none_stays_unvalidated() -> None: + assert _BboxProbe(bbox=None).to_wire() == {"bbox": None} + + +@pytest.mark.parametrize( + ("kwargs", "message"), + [ + ({"lat": 40.0}, "lat and lon must be supplied together"), + ({"lon": -74.0}, "lat and lon must be supplied together"), + ({"radius": "10km"}, "radius needs lat and lon"), + ({"lat": 40.0, "lon": -74.0, "radius": "wide"}, "must be a number optionally suffixed"), + ({"lat": 40.0, "lon": -74.0, "radius": "1001km"}, "greater than 0"), + ({"geo": ["30.27,-97.74"] * 11}, "11 geo shapes"), + ( + {"lat": 40.0, "lon": -74.0, "geo": ["30.27,-97.74"] * 5, "bbox": ["40.4,-74.3,41.0,-73.7"] * 5}, + "11 geo shapes", + ), + ], +) +def test_discover_params_rejects_geo_combinations_the_platform_would_reject( + kwargs: dict[str, object], message: str +) -> None: + with pytest.raises(pydantic.ValidationError, match=message): + DiscoverParams.model_validate(kwargs) + + +def test_discover_params_accepts_ten_shapes_and_a_centre_with_radius() -> None: + DiscoverParams(lat=40.0, lon=-74.0, radius="30mi", geo=["30.27,-97.74"] * 4, bbox=["40.4,-74.3,41.0,-73.7"] * 5) + + +def test_count_params_rejects_lat_without_lon() -> None: + with pytest.raises(pydantic.ValidationError, match="supplied together"): + CountParams(lat=40.0) diff --git a/packages/discolike/tests/test_oauth.py b/packages/discolike/tests/test_oauth.py index d3ad1c0..efc4d6e 100644 --- a/packages/discolike/tests/test_oauth.py +++ b/packages/discolike/tests/test_oauth.py @@ -142,6 +142,7 @@ def handler(request: httpx2.Request) -> httpx2.Response: assert credential.refresh_token == "rt" assert credential.client_id == "client-abc" assert credential.token_endpoint == METADATA.token_endpoint + assert credential.resource == BASE_URL assert int(before) + 3600 <= credential.expires_at <= time.time() + 3600 @@ -293,3 +294,49 @@ def test_credential_config_roundtrip_and_expiry() -> None: assert OAuthCredential.from_config(credential.to_config()) == credential assert credential.expires_within(60, now=950.0) assert not credential.expires_within(60, now=900.0) + + +def test_refresh_sends_resource_and_keeps_it_on_rotated_credential() -> None: + """Without `resource` a server with a default resource may re-bind the refreshed token elsewhere.""" + credential = OAuthCredential( + access_token="old", + refresh_token="rt-old", + expires_at=0.0, + client_id="c", + token_endpoint=METADATA.token_endpoint, + resource=BASE_URL, + ) + seen: list[httpx2.Request] = [] + + def rotating(request: httpx2.Request) -> httpx2.Response: + seen.append(request) + return httpx2.Response(200, json={"access_token": "new", "refresh_token": "rt-new", "expires_in": 60}) + + rotated = refresh(credential, transport=transport_for(rotating)) + assert form(seen[0]) == { + "grant_type": "refresh_token", + "refresh_token": "rt-old", + "client_id": "c", + "resource": BASE_URL, + } + assert rotated.resource == BASE_URL + + +async def test_refresh_async_sends_resource_and_keeps_it_on_rotated_credential() -> None: + credential = OAuthCredential( + access_token="old", + refresh_token="rt-old", + expires_at=0.0, + client_id="c", + token_endpoint=METADATA.token_endpoint, + resource=BASE_URL, + ) + seen: list[httpx2.Request] = [] + + def rotating(request: httpx2.Request) -> httpx2.Response: + seen.append(request) + return httpx2.Response(200, json={"access_token": "new", "refresh_token": "rt-new", "expires_in": 60}) + + rotated = await refresh_async(credential, transport=transport_for(rotating)) + assert form(seen[0])["resource"] == BASE_URL + assert rotated.resource == BASE_URL diff --git a/packages/discolike/tests/test_package.py b/packages/discolike/tests/test_package.py index 69dd262..5700221 100644 --- a/packages/discolike/tests/test_package.py +++ b/packages/discolike/tests/test_package.py @@ -2,4 +2,4 @@ def test_version() -> None: - assert discolike.__version__ == "0.3.2" + assert discolike.__version__ == "0.4.0" diff --git a/packages/discolike/tests/test_requests_contract.py b/packages/discolike/tests/test_requests_contract.py new file mode 100644 index 0000000..3b3a774 --- /dev/null +++ b/packages/discolike/tests/test_requests_contract.py @@ -0,0 +1,138 @@ +"""The generated request models must carry the sub-industry and geo filters.""" + +from discolike.requests import CountParams +from discolike.requests import DiscoverParams +from discolike.resources.companies import CompanyProfile + +NEW_FILTERS = ("sub_industry", "negate_sub_industry", "lat", "lon", "radius", "geo", "bbox") +NEW_OUTPUTS = ("sub_industry", "latitude", "longitude", "geo_precision") + + +SUB_INDUSTRY = ["ROOFING", "CONSTRUCTION/ROOFING"] +NEGATE_SUB_INDUSTRY = ["FOUNDRIES"] +LAT = 40.7128 +LON = -74.006 +RADIUS = "30mi" +BBOX = "40.4,-74.3,41.0,-73.7" +BBOXES = ["40.4,-74.3,41.0,-73.7", "51.2,-0.5,51.7,0.3"] +GEO = ["30.27,-97.74,10km", "52.52,13.405"] + + +class TestGeneratedRequests: + def test_discover_params_carry_the_new_filters(self) -> None: + assert set(NEW_FILTERS) <= set(DiscoverParams.model_fields) + + def test_count_params_carry_the_new_filters(self) -> None: + assert set(NEW_FILTERS) <= set(CountParams.model_fields) + + def test_discover_params_serialize_the_new_filters_onto_the_wire(self) -> None: + wire = DiscoverParams( + sub_industry=SUB_INDUSTRY, + negate_sub_industry=NEGATE_SUB_INDUSTRY, + lat=LAT, + lon=LON, + radius=RADIUS, + ).to_wire() + assert wire["sub_industry"] == ["ROOFING", "CONSTRUCTION/ROOFING"] + assert wire["negate_sub_industry"] == ["FOUNDRIES"] + assert wire["lat"] == 40.7128 + assert wire["lon"] == -74.006 + assert wire["radius"] == "30mi" + assert isinstance(wire["sub_industry"], list) + assert isinstance(wire["lat"], float) + assert isinstance(wire["lon"], float) + assert isinstance(wire["radius"], str) + + def test_count_params_serialize_the_new_filters_onto_the_wire(self) -> None: + wire = CountParams( + sub_industry=SUB_INDUSTRY, + negate_sub_industry=NEGATE_SUB_INDUSTRY, + lat=LAT, + lon=LON, + radius=RADIUS, + ).to_wire() + assert wire["sub_industry"] == ["ROOFING", "CONSTRUCTION/ROOFING"] + assert wire["negate_sub_industry"] == ["FOUNDRIES"] + assert wire["lat"] == 40.7128 + assert wire["lon"] == -74.006 + assert wire["radius"] == "30mi" + + def test_discover_params_serialize_bbox_onto_the_wire(self) -> None: + assert DiscoverParams(bbox=BBOX).to_wire()["bbox"] == [BBOX] + + def test_count_params_serialize_bbox_onto_the_wire(self) -> None: + assert CountParams(bbox=BBOX).to_wire()["bbox"] == [BBOX] + + def test_discover_params_serialize_several_boxes_onto_the_wire(self) -> None: + assert DiscoverParams(bbox=BBOXES).to_wire()["bbox"] == BBOXES + + def test_count_params_serialize_several_boxes_onto_the_wire(self) -> None: + assert CountParams(bbox=BBOXES).to_wire()["bbox"] == BBOXES + + def test_discover_params_serialize_geo_onto_the_wire(self) -> None: + assert DiscoverParams(geo=GEO).to_wire()["geo"] == GEO + + def test_count_params_serialize_geo_onto_the_wire(self) -> None: + assert CountParams(geo=GEO).to_wire()["geo"] == GEO + + def test_a_single_geo_circle_may_be_a_bare_string(self) -> None: + assert DiscoverParams(geo="30.27,-97.74").to_wire()["geo"] == ["30.27,-97.74"] + + def test_geo_and_bbox_and_a_lat_lon_centre_combine(self) -> None: + wire = DiscoverParams(geo=GEO, bbox=BBOXES, lat=LAT, lon=LON, radius=RADIUS).to_wire() + assert wire["geo"] == GEO + assert wire["bbox"] == BBOXES + assert wire["lat"] == LAT + + def test_unset_new_filters_are_dropped_by_exclude_unset(self) -> None: + wire = DiscoverParams(icp_prompt="widgets").to_wire() + for field in NEW_FILTERS: + assert field not in wire + + def test_sub_industry_accepts_a_bare_label(self) -> None: + request = DiscoverParams(sub_industry=["ROOFING"]) + assert request.to_wire()["sub_industry"] == ["ROOFING"] + + def test_sub_industry_accepts_a_parent_qualified_key(self) -> None: + request = DiscoverParams(sub_industry=["CONSTRUCTION/ROOFING"]) + assert request.to_wire()["sub_industry"] == ["CONSTRUCTION/ROOFING"] + + +class TestCompanyProfile: + def test_new_output_fields_exist(self) -> None: + assert set(NEW_OUTPUTS) <= set(CompanyProfile.model_fields) + + def test_sub_industry_defaults_to_none_matching_the_platform_model(self) -> None: + assert CompanyProfile(domain="acme.com").sub_industry is None + + def test_an_empty_mapping_is_preserved(self) -> None: + assert CompanyProfile(domain="acme.com", sub_industry={}).sub_industry == {} + + def test_a_populated_response_parses_the_new_fields(self) -> None: + profile = CompanyProfile.model_validate( + { + "domain": "acme.com", + "sub_industry": {"ROOFING": 0.92, "FOUNDRIES": 0.11}, + "latitude": 40.7128, + "longitude": -74.006, + "geo_precision": "city", + } + ) + assert profile.sub_industry == {"ROOFING": 0.92, "FOUNDRIES": 0.11} + assert isinstance(profile.latitude, float) + assert isinstance(profile.longitude, float) + assert profile.latitude == 40.7128 + assert profile.longitude == -74.006 + assert profile.geo_precision == "city" + + def test_the_pre_rename_coordinate_keys_still_parse(self) -> None: + profile = CompanyProfile.model_validate({"domain": "acme.com", "lat": 40.7128, "lon": -74.006}) + + assert profile.latitude == 40.7128 + assert profile.longitude == -74.006 + + def test_coordinates_serialize_under_their_full_names(self) -> None: + profile = CompanyProfile.model_validate({"domain": "acme.com", "lat": 40.7128, "lon": -74.006}) + + assert profile.to_dict()["latitude"] == 40.7128 + assert profile.to_dict()["longitude"] == -74.006 diff --git a/scripts/gen_requests.py b/scripts/gen_requests.py index 26b8fcd..43e5f14 100644 --- a/scripts/gen_requests.py +++ b/scripts/gen_requests.py @@ -60,6 +60,101 @@ ] +_SUB_INDUSTRY_DESCRIPTION = ( + "Filter by sub-industry, a second-level label scoped to an industry category. Accepts a bare label (ROOFING) " + "or a parent-qualified key (CONSTRUCTION/ROOFING), case-insensitive, up to 50 values. A bare label whose " + "parent category is unambiguous adds that parent to the category filter. Call list-industry-categories for " + "the label list." +) +_RADIUS_DESCRIPTION = ( + "Search radius around lat/lon: a number optionally suffixed with km or mi (50km, 30mi, 50). A bare number is " + "kilometres. Defaults to 50km when lat/lon are supplied, maximum 1000km." +) +_SHAPE_UNION_SENTENCE = ( + "Repeatable: every geo circle, every bbox and the lat/lon/radius centre are OR'd together, up to 10 " + "shapes in total." +) +_GEO_DESCRIPTION = ( + "A circular area to search, written lat,lon or lat,lon,radius (30.27,-97.74 or 30.27,-97.74,30mi). The " + "radius is a number optionally suffixed with km or mi, a bare number meaning kilometres; it defaults to " + "50km and may not exceed 1000km. " + _SHAPE_UNION_SENTENCE +) +_BBOX_DESCRIPTION = ( + "Bounding box as min_lat,min_lon,max_lat,max_lon. Longitudes may wrap the antimeridian (min_lon above " + "max_lon). " + _SHAPE_UNION_SENTENCE +) +_GEO_PROPERTIES: dict[str, dict[str, Any]] = { + "lat": { + "type": "number", + "minimum": -90.0, + "maximum": 90.0, + "nullable": True, + "description": "Latitude of the search centre. Must be supplied together with lon.", + "title": "Lat", + }, + "lon": { + "type": "number", + "minimum": -180.0, + "maximum": 180.0, + "nullable": True, + "description": "Longitude of the search centre. Must be supplied together with lat.", + "title": "Lon", + }, + "radius": {"type": "string", "nullable": True, "description": _RADIUS_DESCRIPTION, "title": "Radius"}, + "geo": { + "type": "array", + "items": {"type": "string"}, + "nullable": True, + "description": _GEO_DESCRIPTION, + "title": "Geo", + }, + "bbox": { + "type": "array", + "items": {"type": "string"}, + "nullable": True, + "description": _BBOX_DESCRIPTION, + "title": "Bbox", + }, +} +_SUB_INDUSTRY_PROPERTIES: dict[str, dict[str, Any]] = { + "sub_industry": { + "type": "array", + "items": {"type": "string"}, + "maxItems": 50, + "nullable": True, + "description": _SUB_INDUSTRY_DESCRIPTION, + "title": "Sub Industry", + }, + "negate_sub_industry": { + "type": "array", + "items": {"type": "string"}, + "maxItems": 50, + "nullable": True, + "description": ( + "Exclude specified sub-industries. Same format as sub_industry; does not affect the category filter." + ), + "title": "Negate Sub Industry", + }, +} + +# Properties the SDK ships before the deployed spec has them. Merged in only while the spec +# lacks them, so each entry clears itself once the platform release lands -- generation prints +# the ones that have, to be deleted here. +PENDING_PROPERTIES: dict[str, dict[str, dict[str, Any]]] = { + "DiscoverParams": _GEO_PROPERTIES, + "CountParams": _GEO_PROPERTIES, +} + +# Properties generated from this schema rather than the spec's, whatever the spec says. The +# platform's sub-industry enum lists parent-qualified keys only; a bare label reaches it through +# a server-side normalizer with no client-side counterpart, so generating that enum would reject +# values the API accepts. +PROPERTY_OVERRIDES: dict[str, dict[str, dict[str, Any]]] = { + "DiscoverParams": _SUB_INDUSTRY_PROPERTIES, + "CountParams": _SUB_INDUSTRY_PROPERTIES, +} + + @dataclass(frozen=True) class Route: class_name: str @@ -191,8 +286,25 @@ def normalize_schema(schema: dict[str, Any]) -> dict[str, Any]: return {key: value for key, value in schema.items() if key != "additionalProperties"} | {"properties": properties} +def apply_overlays(*, kept: dict[str, dict[str, Any]]) -> dict[str, dict[str, Any]]: + for name, properties in PENDING_PROPERTIES.items(): + if (schema := kept.get(name)) is None: + continue + existing = schema.setdefault("properties", {}) + for prop, prop_schema in properties.items(): + if prop in existing: + print(f"note: the spec now has {name}.{prop}; drop it from PENDING_PROPERTIES") + continue + existing[prop] = copy.deepcopy(prop_schema) + for name, properties in PROPERTY_OVERRIDES.items(): + if (schema := kept.get(name)) is None: + continue + schema.setdefault("properties", {}).update(copy.deepcopy(properties)) + return kept + + def build_codegen_spec(*, spec: dict[str, Any], routes: list[Route]) -> dict[str, Any]: - kept = prune(spec=spec, requested=request_schemas(spec=spec, routes=routes)) + kept = apply_overlays(kept=prune(spec=spec, requested=request_schemas(spec=spec, routes=routes))) return { "openapi": "3.1.0", "info": {"title": "discolike request models", "version": "0"}, diff --git a/uv.lock b/uv.lock index efe280e..cb4f963 100644 --- a/uv.lock +++ b/uv.lock @@ -3,8 +3,8 @@ revision = 3 requires-python = ">=3.10" resolution-markers = [ "python_full_version >= '3.14' and sys_platform == 'emscripten'", - "python_full_version >= '3.12' and python_full_version < '3.14' and sys_platform == 'emscripten'", "python_full_version >= '3.14' and sys_platform != 'emscripten'", + "python_full_version >= '3.12' and python_full_version < '3.14' and sys_platform == 'emscripten'", "(python_full_version < '3.14' and sys_platform != 'emscripten') or (python_full_version < '3.12' and sys_platform == 'emscripten')", ] @@ -374,7 +374,7 @@ dev = [ [[package]] name = "discolike-cli" -version = "0.3.2" +version = "0.4.0" source = { editable = "packages/discolike-cli" } dependencies = [ { name = "discolike" },