from __future__ import annotations import fnmatch import html import json import os import re from functools import lru_cache from html.parser import HTMLParser from pathlib import Path from urllib.parse import parse_qs, urlparse _MAX_READ_CHARS = 4000 _MAX_GREP_MATCHES = 20 _MAX_GLOB_MATCHES = 50 _MAX_ORACLE_PAGE_CHARS = 24000 _MAX_ORACLE_CONTEXT_CHARS = 80000 def _normalize_data_dirs(data_dirs: list[str] | tuple[str, ...] | str | None, project_root: Path) -> list[str]: if data_dirs is None: return [] if isinstance(data_dirs, str): items = [part.strip() for chunk in data_dirs.split(os.pathsep) for part in chunk.split(",")] else: items = [str(item).strip() for item in data_dirs] resolved: list[str] = [] for item in items: if not item: continue path = Path(item).expanduser() if not path.is_absolute(): path = project_root / path resolved.append(str(path)) return resolved def resolve_docs_roots(data_dirs: list[str] | tuple[str, ...] | str | None = None) -> list[str]: project_root = Path(__file__).resolve().parents[3] env_value = os.environ.get("OFFICEQA_DOCS_DIR", "").strip() candidates = _normalize_data_dirs(data_dirs, project_root) candidates.extend(_normalize_data_dirs(env_value, project_root)) candidates.extend([ str(project_root / "data" / "officeqa_docs_official"), str(project_root / "data" / "officeqa_smoke_docs"), os.path.expanduser("~/officeqa-sparse/treasury_bulletins_parsed"), os.path.expanduser("~/officeqa/treasury_bulletins_parsed"), ]) roots: list[str] = [] seen: set[str] = set() for candidate in candidates: path = Path(candidate).expanduser() if not path.is_dir(): continue transformed = path / "transformed" resolved = str((transformed if transformed.is_dir() else path).resolve()) if resolved in seen: continue seen.add(resolved) roots.append(resolved) if not roots: raise FileNotFoundError("OfficeQA docs directory not found. Set OFFICEQA_DOCS_DIR or env.data_dirs.") return roots def _is_allowed(path: str, allowed_roots: list[str], allowed_files: list[str]) -> bool: try: resolved = str(Path(path).resolve()) except FileNotFoundError: return False if not any(resolved.startswith(root + os.sep) or resolved == root for root in allowed_roots): return False if not allowed_files: return True base = os.path.basename(resolved) return base in allowed_files def resolve_candidate_files(source_files: list[str], allowed_roots: list[str]) -> list[str]: resolved: list[str] = [] seen: set[str] = set() for root in allowed_roots: for dirpath, _, filenames in os.walk(root): for filename in filenames: if source_files and filename not in source_files: continue full = str(Path(dirpath, filename).resolve()) if full in seen: continue seen.add(full) resolved.append(full) return resolved def _as_list(value: object) -> list[str]: if value is None: return [] if isinstance(value, list): return [str(item).strip() for item in value if str(item).strip()] text = str(value).strip() if not text: return [] try: loaded = json.loads(text) except json.JSONDecodeError: loaded = None if isinstance(loaded, list): return [str(item).strip() for item in loaded if str(item).strip()] if "\n" in text: return [part.strip() for part in text.splitlines() if part.strip()] return [text] def _extract_page_number(source_doc: str) -> int | None: text = str(source_doc or "").strip() if not text: return None parsed = urlparse(text) query = parse_qs(parsed.query) for key in ("page", "pagenum", "page_id"): for raw_value in query.get(key, []): try: return int(str(raw_value).strip()) except ValueError: continue match = re.search(r"(?:[?&]|^)page=(\d+)", text) if match: return int(match.group(1)) return None def _iter_oracle_refs(source_files: object, source_docs: object) -> list[tuple[str, int, str]]: files = _as_list(source_files) docs = _as_list(source_docs) refs: list[tuple[str, int, str]] = [] seen: set[tuple[str, int, str]] = set() if not files or not docs: return refs for index, source_doc in enumerate(docs): page_number = _extract_page_number(source_doc) if page_number is None: continue if index < len(files): source_file = files[index] elif len(files) == 1: source_file = files[0] else: continue key = (source_file, page_number, source_doc) if key in seen: continue seen.add(key) refs.append(key) return refs def _parsed_root_candidates(docs_roots: list[str]) -> list[Path]: candidates: list[Path] = [] seen: set[str] = set() for root in docs_roots: path = Path(root).expanduser() for candidate in ( path, path.parent, path / "treasury_bulletins_parsed", path.parent / "treasury_bulletins_parsed", ): resolved = str(candidate.resolve()) if candidate.exists() else str(candidate) if resolved in seen: continue seen.add(resolved) candidates.append(candidate) return candidates def _locate_parsed_json(source_file: str, docs_roots: list[str]) -> Path | None: source_path = Path(str(source_file).strip()) stem = source_path.stem if source_path.suffix else source_path.name if not stem: return None candidate_names = [stem + ".json"] if source_path.suffix == ".json": candidate_names.insert(0, source_path.name) for root in _parsed_root_candidates(docs_roots): for name in candidate_names: path = root / "jsons" / name if path.is_file(): return path return None class _TableMarkdownParser(HTMLParser): def __init__(self) -> None: super().__init__(convert_charrefs=True) self.rows: list[list[str]] = [] self._row: list[str] | None = None self._cell: list[str] | None = None def handle_starttag(self, tag: str, attrs: list[tuple[str, str | None]]) -> None: if tag.lower() == "tr": self._row = [] elif tag.lower() in {"td", "th"} and self._row is not None: self._cell = [] def handle_data(self, data: str) -> None: if self._cell is not None: self._cell.append(data) def handle_endtag(self, tag: str) -> None: normalized_tag = tag.lower() if normalized_tag in {"td", "th"} and self._cell is not None and self._row is not None: cell = re.sub(r"\s+", " ", "".join(self._cell)).strip() self._row.append(cell) self._cell = None elif normalized_tag == "tr" and self._row is not None: if any(cell for cell in self._row): self.rows.append(self._row) self._row = None self._cell = None def _escape_markdown_cell(value: str) -> str: return str(value).replace("\n", " ").replace("|", "\\|").strip() def _html_table_to_markdown(raw_html: str) -> str: parser = _TableMarkdownParser() try: parser.feed(raw_html) except Exception: # noqa: BLE001 parser.rows = [] rows = parser.rows if not rows: text = re.sub(r"(?is)<[^>]+>", " ", raw_html) return re.sub(r"\s+", " ", html.unescape(text)).strip() width = max(len(row) for row in rows) normalized_rows = [row + [""] * (width - len(row)) for row in rows] header = normalized_rows[0] body = normalized_rows[1:] lines = [ "| " + " | ".join(_escape_markdown_cell(cell) for cell in header) + " |", "| " + " | ".join(["---"] * width) + " |", ] lines.extend("| " + " | ".join(_escape_markdown_cell(cell) for cell in row) + " |" for row in body) return "\n".join(lines) def _render_parsed_content(content: str) -> str: text = content.strip() if not text: return "" if " set[int]: page_ids: set[int] = set() bbox = element.get("bbox") if not isinstance(bbox, list): return page_ids for box in bbox: if not isinstance(box, dict): continue raw_page_id = box.get("page_id") try: page_ids.add(int(raw_page_id)) except (TypeError, ValueError): continue return page_ids @lru_cache(maxsize=256) def _load_parsed_elements(json_path: str) -> tuple[dict, ...]: with open(json_path, encoding="utf-8") as f: payload = json.load(f) document = payload.get("document") if isinstance(payload, dict) else {} elements = document.get("elements") if isinstance(document, dict) else [] if not isinstance(elements, list): return () return tuple(element for element in elements if isinstance(element, dict)) @lru_cache(maxsize=2048) def _render_parsed_page(json_path: str, page_number: int) -> str: rendered: list[str] = [] for element in _load_parsed_elements(json_path): if page_number not in _element_page_ids(element): continue content = element.get("content") if not isinstance(content, str) or not content.strip(): continue section = _render_parsed_content(content) if section: rendered.append(section) return "\n\n".join(rendered).strip() def build_oracle_parsed_pages_context( source_files: object, source_docs: object, docs_roots: list[str], *, max_page_chars: int = _MAX_ORACLE_PAGE_CHARS, max_total_chars: int = _MAX_ORACLE_CONTEXT_CHARS, evidence_note: str = "Treat it as primary document evidence and combine it with custom web search results when useful.", ) -> str: """Render oracle parsed OfficeQA pages referenced by source_docs/source_files.""" refs = _iter_oracle_refs(source_files, source_docs) if not refs: return "" blocks: list[str] = [] total_chars = 0 seen_pages: set[tuple[str, int]] = set() for source_file, page_number, source_doc in refs: json_path = _locate_parsed_json(source_file, docs_roots) if json_path is None: continue page_key = (str(json_path), page_number) if page_key in seen_pages: continue seen_pages.add(page_key) page_text = _render_parsed_page(str(json_path), page_number) if not page_text: continue if len(page_text) > max_page_chars: omitted = len(page_text) - max_page_chars page_text = page_text[:max_page_chars].rstrip() + f"\n\n[... {omitted} characters omitted from this parsed page ...]" block = ( f"### {source_file} page {page_number}\n" f"Source URL: {source_doc}\n\n" f"{page_text}" ) if total_chars + len(block) > max_total_chars: remaining = max_total_chars - total_chars if remaining <= 0: break block = block[:remaining].rstrip() + "\n\n[... oracle parsed page context truncated ...]" blocks.append(block) break blocks.append(block) total_chars += len(block) if not blocks: return "" return ( "The following content is pre-parsed from the oracle OfficeQA source page(s). " f"{evidence_note.strip()}\n\n" + "\n\n".join(blocks) ) def run_tool(name: str, arguments: dict, *, allowed_roots: list[str], allowed_files: list[str]) -> tuple[str, str]: if name == "glob": pattern = str(arguments.get("pattern") or "*") matches: list[str] = [] for root in allowed_roots: for dirpath, _, filenames in os.walk(root): for filename in filenames: if allowed_files and filename not in allowed_files: continue rel = os.path.relpath(os.path.join(dirpath, filename), root) if fnmatch.fnmatch(rel, pattern) or fnmatch.fnmatch(filename, pattern): matches.append(os.path.join(dirpath, filename)) if len(matches) >= _MAX_GLOB_MATCHES: break if len(matches) >= _MAX_GLOB_MATCHES: break return f"glob(pattern={pattern!r})", "\n".join(matches) if matches else "[no matches]" if name == "read": path = str(arguments.get("path") or "") if not path: return "read(path='')", "[read error: missing path]" if not _is_allowed(path, allowed_roots, allowed_files): return f"read(path={path!r})", "[read error: path not allowed]" start = max(int(arguments.get("start") or 1), 1) limit = max(int(arguments.get("limit") or 80), 1) with open(path, encoding="utf-8") as f: lines = f.readlines() excerpt = "".join(lines[start - 1:start - 1 + limit]) return f"read(path={path!r}, start={start}, limit={limit})", excerpt[:_MAX_READ_CHARS] or "[empty file]" if name == "grep": pattern = str(arguments.get("pattern") or "").lower() path = str(arguments.get("path") or "") if not pattern or not path: return f"grep(pattern={pattern!r}, path={path!r})", "[grep error: missing pattern or path]" if not _is_allowed(path, allowed_roots, allowed_files): return f"grep(pattern={pattern!r}, path={path!r})", "[grep error: path not allowed]" matches: list[str] = [] with open(path, encoding="utf-8") as f: for idx, line in enumerate(f, start=1): if pattern in line.lower(): matches.append(f"{idx}: {line.rstrip()}") if len(matches) >= _MAX_GREP_MATCHES: break return f"grep(pattern={pattern!r}, path={path!r})", "\n".join(matches) if matches else "[no matches]" return name, f"[tool error: unknown tool {name}]"