diff --git a/.github/workflows/build_pr.yml b/.github/workflows/build_pr.yml index de063bb..1600a06 100644 --- a/.github/workflows/build_pr.yml +++ b/.github/workflows/build_pr.yml @@ -1,4 +1,4 @@ -name: Build and publish release on new tag +name: Build and publish PR image on: pull_request: types: [opened, synchronize, reopened] @@ -20,19 +20,11 @@ jobs: - name: Set up Docker Buildx id: buildx uses: docker/setup-buildx-action@v3 - - name: Build and push SAIST-Lite - run: | - docker buildx build . \ - --push \ - --progress=plain \ - --tag docker.io/punksecurity/saist:lite-pr-${{ github.event.number }} \ - --platform linux/amd64 \ - --target saist - - name: Build and push SAIST - run: | - docker buildx build . \ - --push \ - --progress=plain \ - --tag docker.io/punksecurity/saist:pr-${{ github.event.number }} \ - --platform linux/amd64 \ - --target saist-tex + - name: Build and push SAIST + run: | + docker buildx build . \ + --push \ + --progress=plain \ + --tag docker.io/punksecurity/saist:pr-${{ github.event.number }} \ + --platform linux/amd64 \ + --target saist diff --git a/.github/workflows/build_release.yml b/.github/workflows/build_release.yml index e8079f7..965dd23 100644 --- a/.github/workflows/build_release.yml +++ b/.github/workflows/build_release.yml @@ -28,4 +28,4 @@ jobs: VERSION: ${{ steps.version.outputs.version }} with: push: true - targets: "full,lite" + targets: "release" diff --git a/.github/workflows/tests.yml b/.github/workflows/tests.yml new file mode 100644 index 0000000..1b05748 --- /dev/null +++ b/.github/workflows/tests.yml @@ -0,0 +1,30 @@ +name: Tests + +on: + pull_request: + types: [opened, synchronize, reopened] + +jobs: + pytest: + runs-on: ubuntu-latest + + steps: + - name: Check out code + uses: actions/checkout@v4 + + - name: Set up Python + uses: actions/setup-python@v5 + with: + python-version: "3.12" + cache: pip + cache-dependency-path: | + requirements.txt + requirements-dev.txt + + - name: Install dependencies + run: | + python -m pip install --upgrade pip + pip install -r requirements-dev.txt + + - name: Run pytest + run: pytest -q diff --git a/Dockerfile b/Dockerfile index 02d1193..e64c03b 100644 --- a/Dockerfile +++ b/Dockerfile @@ -19,43 +19,8 @@ COPY saist . # Exports ENV SAIST_COMMAND "docker run punksecurity/saist" ENV SAIST_CSV_PATH "/app/results.csv" -ENV SAIST_TEX_FILENAME "report.tex" ENV SAIST_PDF_FILENAME "report.pdf" ENV SAIST_WEB_HOST "0.0.0.0" ENV PYTHONUNBUFFERED 1 ENTRYPOINT [ "python3", "/app/main.py" ] CMD [ "-h" ] - -FROM alpine as tex-dl -WORKDIR /tmp -RUN mkdir -p /opt/texlive/bin - -RUN wget https://ftp.math.utah.edu/pub/texlive-utah/bin/aarch64-alpine322.tar.xz && \ - tar xvf aarch64-alpine322.tar.xz && ls -ltra && \ - mv aarch64-alpine322 /opt/texlive/bin/aarch64-linuxmusl - -RUN wget https://ftp.math.utah.edu/pub/texlive-utah/bin/x86_64-alpine322.tar.xz && \ - tar xvf x86_64-alpine322.tar.xz && \ - mv x86_64-alpine322 /opt/texlive/bin/x86_64-linuxmusl - -FROM saist AS saist-tex - -ARG TL_MIRROR="https://texlive.info/CTAN/systems/texlive/tlnet" - -COPY saist/latex/texlive.profile /tmp -COPY --from=tex-dl /opt/texlive/bin /tmp/texlive/bin -RUN apk add --no-cache perl curl fontconfig xz && \ - mkdir -p "/tmp/texlive" && cd "/tmp/texlive" && \ - wget "$TL_MIRROR/install-tl-unx.tar.gz" && \ - tar xzvf ./install-tl-unx.tar.gz && \ - "./install-tl-"*"/install-tl" --location "$TL_MIRROR" --custom-bin=/tmp/texlive/bin/$(uname -m)-linuxmusl -profile "/tmp/texlive.profile" && \ - rm -vf "/opt/texlive/install-tl" && \ - rm -vf "/opt/texlive/install-tl.log" && \ - rm -vrf /tmp/* - -ENV PATH="${PATH}:/opt/texlive/bin/custom" - -ARG TL_PACKAGES="lineno titlesec upquote minted blindtext booktabs fontawesome latexmk parskip xcolor" - -RUN tlmgr update --self && \ - tlmgr install ${TL_PACKAGES} diff --git a/README.md b/README.md index 201db28..f820a55 100644 --- a/README.md +++ b/README.md @@ -29,10 +29,12 @@ We support OLLAMA for local / offline code scanning. - **Diff scanning**: Git commits, branches, or PRs - **Multi-LLM support**: `OpenAI`, `Anthropic`, `Bedrock`, `DeepSeek`, `Gemini`, `Ollama` - **Filesystem, Git, GitHub PR scanning modes** -- **Pattern-based file inclusion/exclusion** using `.saist.include` and `.saist.ignore` -- **Interactive chat** with your findings -- **Web server** UI to view results -- **CSV export** of findings +- **Pattern-based file inclusion/exclusion** using `.saist.include` and `.saist.ignore` +- **Project-specific analysis skills** loaded from Markdown files to teach SAIST app routing, authorization, framework conventions, and other local security context +- **LLM-generated analysis skills** for bootstrapping those files in a separate run +- **Interactive chat** with your findings +- **Web server** UI to view results +- **CSV export** of findings - **PDF report**: Generate PDF reports of SAIST findings - **CI/CD pipeline friendly** (exit 1 on findings) @@ -73,16 +75,18 @@ export SAIST_LLM_API_KEY=your-api-key | Task | Command | |:-----|:--------| -| Get a DevSecOps poem | `saist/main.py --llm openai poem` | -| Scan a local folder | `saist/main.py --llm deepseek filesystem /path/to/code` | -| Scan a local folder with ollama from within docker| `docker run --network=host -v :/vulnerableapp -v $PWD/reporting:/app/reporting punksecurity/saist --llm ollama --llm-model gemma3:4b fileystem /vulnerableapp` | +| Get a DevSecOps poem | `saist/main.py --llm openai poem` | +| Scan a local folder | `saist/main.py --llm deepseek filesystem /path/to/code` | +| Scan a local folder file-by-file | `saist/main.py --llm deepseek --deep filesystem /path/to/code` | +| Scan a local folder with ollama from within docker| `docker run --network=host -v :/vulnerableapp -v $PWD/reporting:/app/reporting punksecurity/saist --llm ollama --llm-model gemma3:4b fileystem /vulnerableapp` | | Scan a local Git repo | `saist/main.py --llm openai git /path/to/repo` | | Scan a local Git repo (branch diff) | `saist/main.py --llm openai git /path/to/repo --ref-for-compare main --ref-to-compare feature-branch` | | Scan a GitHub PR (and update the PR) | `saist/main.py --llm anthropic github yourorg/yourrepo 1234 --github-token your-token` | | Launch web server to view findings | `saist/main.py --llm deepseek --web filesystem /path/to/code` | | Interactive shell after scanning | `saist/main.py --llm ollama --interactive filesystem /path/to/code` | -| Export findings as CSV | `saist/main.py --llm openai --csv filesystem /path/to/code` | -| Scan with docker and export findings as PDF report | `docker run -v :/vulnerableapp -v $PWD/reporting:/app/reporting punksecurity/saist --llm openai --pdf filesystem /vulnerableapp` | +| Export findings as CSV | `saist/main.py --llm openai --csv filesystem /path/to/code` | +| Generate analysis skills | `saist/main.py --llm openai --generate-skills filesystem /path/to/code` | +| Scan with docker and export findings as PDF report | `docker run -v :/vulnerableapp -v $PWD/reporting:/app/reporting punksecurity/saist --llm openai --pdf filesystem /vulnerableapp` | | Scan with docker and export findings as PDF report with a project title | `docker run -v :/vulnerableapp -v $PWD/reporting:/app/reporting punksecurity/saist --llm openai --pdf --project-name "Project Name" filesystem /vulnerableapp` | | Scan with docker and retain cache for future runs | `docker run -v :/vulnerableapp -v $PWD/SAISTCache:/app/SAISTCache punksecurity/saist --llm openai filesystem /vulnerableapp` | | Change caching folder | `saist/main.py --llm openai --cache-folder /path/to/cache filesystem /path/to/code` | @@ -106,7 +110,7 @@ saist respects **file include/exclude rules** via two optional files in the root - `build/` will ignore the entire build folder - `*.log` will ignore all log files -You can also provide include/exclude patterns using the command-line arguments `--include` and `--exclude`. +You can also provide include/exclude patterns using the command-line arguments `--include` and `--exclude`. - Patterns provided via command-line arguments are appended to any patterns loaded from the rule files. - Examples: - `--include '**/*.py' --include '**/*.ts'` includes all Python and TypeScript files @@ -132,23 +136,51 @@ docs/ This setup will: - Only scan `.py`, `.ts`, and specific `.js` files - Ignore anything under `tests/` and `docs/` ---- - - -## ๐Ÿ“„ PDF report generation +--- + +## ๐Ÿง  Analysis Skills + +SAIST can load project-specific analysis skill files from `.saist/skills/*.md`. These files are added to the security review prompt so future scans understand application-specific details such as routing, authentication, authorization, framework conventions, data access, validation boundaries, dependencies, configuration, and security-sensitive workflows. + +Generate an initial set of skill files as a separate run: + +```bash +saist/main.py --llm openai --generate-skills filesystem /path/to/code +``` + +Then review or edit the generated Markdown files and run SAIST normally. Skill files are loaded automatically on future scans: + +```bash +saist/main.py --llm openai filesystem /path/to/code +``` + +Useful options: + +| Option | Description | +|:------|:------------| +| `--skills-path` | Folder containing skill Markdown files. Defaults to `.saist/skills` under the scanned project. | +| `--generate-skills` | Ask the configured LLM to generate skill files and then exit. | +| `--overwrite-skills` | Replace existing skill files during generation. Without this, existing files are preserved. | +| `--disable-skills` | Do not load skill files during analysis. | +| `--skills-max-bytes` | Limit total skill guidance added to analysis prompts. | +| `--skills-sample-files` / `--skills-sample-bytes` | Control how much project context is sampled when generating skills. | + +When skills are loaded, SAIST salts its findings cache with the skill content so updated guidance gets a fresh analysis run. + +--- + + +## ๐Ÿ“„ PDF report generation saist allows you to generate PDF reports summarizing your findings, making it easier to share insights with your team. -To create a PDF report, simply use the `--pdf` flag when running the scan. By default, the report will be saved to -`reporting/report.pdf`. You can customize the filename by using the `--pdf-filename` option followed by your desired -filename. - -To add a project name onto the title page of the PDF report, use the `--project-name` option followed by your desired title. - -> It is recommended to use the provided Docker image for generating PDF reports, as it includes the necessary TeX suite, -which can be quite large. This ensures that all dependencies are met and the report is generated properly. - -If not, you need to install latexmk to make it work. +To create a PDF report, use the `--pdf` flag when running the scan. By default, the report will be saved to +`reporting/report.pdf`. You can customize the filename by using the `--pdf-filename` option followed by your desired +filename. + +To add a project name onto the title page of the PDF report, use the `--project-name` option followed by your desired title. + +PDF reports are generated with the built-in ReportLab renderer, so no external document-rendering toolchain is required. ### ๐Ÿ‹ Example (Docker) @@ -169,13 +201,23 @@ docker run -v$PWD/code:/code -v$PWD/reporting:/app/reporting punksecurity/saist | Option | Description | |:------|:------------| -| `--llm` | Select LLM (`anthropic`, `deepseek`, `gemini`, `ollama`, `openai`) | -| `--llm-api-key` | API key for your LLM | -| `--llm-model` | (Optional) Specific model (e.g., `gpt-4o`) | -| `--interactive` | Chat with the LLM after scan | -| `--web` | Launch a local web server | -| `--disable-tools` | Disable tool use during file analysis to reduce LLM token usage | -| `--disable-caching` | Disable finding caching during file analysis | +| `--llm` | Select LLM (`anthropic`, `azure-foundry`, `bedrock`, `deepseek`, `gemini`, `ollama`, `openai`) | +| `--llm-api-key` | API key for your LLM | +| `--llm-model` | (Optional) Specific model (e.g., `gpt-4o`) | +| `--thinking` | Pydantic AI thinking effort: `minimal`, `low`, `medium`, `high`, `xhigh`, or `disabled` | +| `--openai-base-uri` | Base URI for OpenAI-compatible services. Can also be set with `SAIST_OPENAI_BASE_URI`. | +| `--azure-openai-endpoint` | Azure AI Foundry or Azure OpenAI endpoint. Can also be set with `AZURE_OPENAI_ENDPOINT`; `/openai/v1/` endpoints use the Responses API without `api-version`. | +| `--azure-openai-api-version` | Azure OpenAI API version for non-v1 endpoints. Can also be set with `OPENAI_API_VERSION`. | +| `--interactive` | Chat with the LLM after scan | +| `--web` | Launch a local web server | +| `--disable-tools` | Disable tool use during file analysis to reduce LLM token usage | +| `--deep` | For filesystem scans, analyze every file individually. Without this, filesystem scans send a file inventory and let the LLM inspect files with tools, then report file coverage. | +| `--iterations` | Number of tool-driven filesystem scan passes to run when `--deep` is not set. Defaults to `1`; concurrency is capped by `--llm-rate-limit`. | +| `--skills-path` | Folder containing SAIST analysis skill Markdown files | +| `--generate-skills` | Generate SAIST analysis skill files and exit | +| `--overwrite-skills` | Replace existing skill files during skill generation | +| `--disable-skills` | Do not load skill files during analysis | +| `--disable-caching` | Disable finding caching during file analysis | | `--skip-line-length-check` | Skip checking files for a maximum line length | | `--max-line-length` | Maximum allowed line length, files with lines longer than this value will be skipped | | `--i, --include` | Pattern to explicitly include | diff --git a/docker-bake.hcl b/docker-bake.hcl index a929d46..e5f7aca 100644 --- a/docker-bake.hcl +++ b/docker-bake.hcl @@ -11,16 +11,11 @@ target "base" { target "nightly" { inherits = ["base"] tags = ["docker.io/punksecurity/saist:nightly"] -} - -target "lite" { - inherits = ["base"] - tags = ["docker.io/punksecurity/saist:lite-${VERSION}", "docker.io/punksecurity/saist:lite-latest"] target = "saist" } -target "full" { +target "release" { inherits = ["base"] tags = ["docker.io/punksecurity/saist:${VERSION}", "docker.io/punksecurity/saist:latest"] - target = "saist-tex" + target = "saist" } diff --git a/pytest.ini b/pytest.ini new file mode 100644 index 0000000..3d0ba62 --- /dev/null +++ b/pytest.ini @@ -0,0 +1,3 @@ +[pytest] +pythonpath = saist +testpaths = tests diff --git a/requirements-dev.txt b/requirements-dev.txt new file mode 100644 index 0000000..31c406a --- /dev/null +++ b/requirements-dev.txt @@ -0,0 +1,2 @@ +-r requirements.txt +pytest==8.3.5 diff --git a/requirements.txt b/requirements.txt index a1eaa77..2ab0fd9 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,10 +1,11 @@ -pydantic==2.11.7 +pydantic==2.12.5 requests==2.32.4 GitPython==3.1.44 python-dotenv==1.1.1 -pydantic-ai==0.3.7 +pydantic-ai-slim[anthropic,bedrock,google,openai]==1.93.0 aiofiles==24.1.0 rich==14.0.0 flask==3.1.1 ollama==0.5.1 gitignore-parser==0.1.13 +reportlab==4.5.0 diff --git a/saist/latex/__init__.py b/saist/latex/__init__.py deleted file mode 100644 index 051f833..0000000 --- a/saist/latex/__init__.py +++ /dev/null @@ -1,88 +0,0 @@ -from dataclasses import dataclass -from models import FindingContext -from llm.adapters import BaseLlmAdapter -from jinja2 import Environment, FileSystemLoader -import subprocess -import logging -import os -import re -import datetime - -logger = logging.getLogger("saist.latex") - - -# TODO: Container 'awareness' -@dataclass -class Latex: - _DEFAULT_TEX_TEMPLATE = "report.tex.jinja" - _DEFAULT_OUTPUT_DIR = "reporting" - - llm: BaseLlmAdapter - project: str - findings: list[FindingContext] - comment: str - - def run(self, args): - pdf_path = os.path.join(self._DEFAULT_OUTPUT_DIR, args.pdf_filename) - tex_path = os.path.join(self._DEFAULT_OUTPUT_DIR, args.tex_filename) - - self._write_tex(tex_path) - logger.info(f"Written TeX file to: '{tex_path}'") - - if args.pdf: - print("\n๐Ÿ“ Generating PDF report...") - - rc = subprocess.run( - ["latexmk", "-pdf", "-f", "-interaction=nonstopmode", f"-outdir={self._DEFAULT_OUTPUT_DIR}", tex_path], stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL, - ).returncode - - if rc != 0: - logger.error(f"Unable to build PDF '{tex_path}' -> '{pdf_path}'") - exit(1) - - print(f"โœจ Written report to '{pdf_path}'\n") - - logger.debug(f"Cleaning up auxiliary files in '{self._DEFAULT_OUTPUT_DIR}'") - subprocess.run(["latexmk", "-c"], cwd=self._DEFAULT_OUTPUT_DIR, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL, ) - - def _render_tex(self) -> str: - env = Environment( - loader=FileSystemLoader(os.path.join(os.path.dirname(__file__), "tex")) - ) - env.globals.update(escape_tex=self._escape_tex) - template = env.get_template(self._DEFAULT_TEX_TEMPLATE) - template.globals['now'] = datetime.datetime.now().isoformat(sep=' ', timespec='minutes') - return template.render({"model": self.llm, "project": self.project, "findings": self.findings, "comment": self.comment}) - - def _write_tex(self, tex_path: str): - try: - os.makedirs(os.path.dirname(tex_path), exist_ok=True) - with open(tex_path, "w") as fp: - fp.write(self._render_tex()) - except Exception as e: - logger.error(f"Unable to write TeX file to '{tex_path}': {e}") - exit(1) - - def _escape_tex(self, text: str) -> str: - if not text: - return "" - - specials = { - "&": r"\&", - "%": r"\%", - "$": r"\$", - "_": r"\_", - "{": r"\{", - "}": r"\}", - "#": r"\#", - "~": r"\textasciitilde{}", - "^": r"\^{}", - "\\": r"\textbackslash{}", - } - - def repl(match: re.Match) -> str: - char = match.group(0) - return specials[char] - - pattern = r'.(?<=[&%$_{}#~^\\])' - return re.sub(pattern, repl, text) diff --git a/saist/latex/tex/brand.tex b/saist/latex/tex/brand.tex deleted file mode 100644 index c9d78f0..0000000 --- a/saist/latex/tex/brand.tex +++ /dev/null @@ -1,22 +0,0 @@ -\def\Primary{111B29} -\def\Secondary{0C2540} - -\def\Blue{5BBFCF} -\def\Purple{A25DCE} -\def\Orange{CE985D} -\def\Green{87CE6D} -\def\Red{CE5D5D} - -\def\LogoTitlepage{latex/tex/image/PunkSecurityLogo.png} - -\usepackage[dvipsnames]{xcolor} -\newcommand{\newcolor}[2]{\definecolor{#1}{HTML}{#2}} - -\newcolor{ColorPrimary}{\Primary} -\newcolor{ColorSecondary}{\Secondary} - -\newcolor{ColorBlue}{\Blue} -\newcolor{ColorPurple}{\Purple} -\newcolor{ColorOrange}{\Orange} -\newcolor{ColorGreen}{\Green} -\newcolor{ColorRed}{\Red} \ No newline at end of file diff --git a/saist/latex/tex/report.tex.jinja b/saist/latex/tex/report.tex.jinja deleted file mode 100644 index 2df1518..0000000 --- a/saist/latex/tex/report.tex.jinja +++ /dev/null @@ -1,67 +0,0 @@ -\input{latex/tex/style.tex} - -\extrafloats{1000} - -\begin{document} - -\begin{titlepage} - \centering - \vspace*{5cm} - \includegraphics[width=0.4\textwidth]{\LogoTitlepage}\\[0.5cm] -{\Huge\sffamily\bfseries\color{ColorSecondary}{AI generated code review}\\[0.75em] -\Large\sffamily created by \href{https://github.com/punk-security/SAIST}{\textbf{SAIST}}}\\[0.75em] -{%- if model.model_vendor is not none -%} -\Large\sffamily\color{ColorSecondary}{{ escape_tex(model.model_vendor) }}: -{%- endif -%} -\Large\sffamily\color{ColorSecondary}{{ escape_tex(model.model_name) }} -{%- if project is not none -%} - \\[0.75em] - \Large\sffamily\color{ColorSecondary}{{ escape_tex(project) }}\\[0.75em] -{%- endif -%} - -\vspace{2.5cm} -\centering\Large\sffamily\color{ColorSecondary}{{ now }} - -\end{titlepage} - -\tableofcontents - -\section{Summary} - -{{ escape_tex(comment) }} - -\newpage -\section{Findings} - -{%- for finding in findings if finding.cwe -%} - {%- set title = escape_tex(finding.title) -%} - {%- set issue = escape_tex(finding.issue) -%} - {%- set file = escape_tex(finding.file) -%} - {%- set recommendation = escape_tex(finding.recommendation) -%} -\subsection{ {{ finding.cwe }} - {{ file }} - {{title}} } - -\textbf{Priority:}\; -{%- if finding.priority > 8 -%}\colorbox{ColorRed}{Critical} -{%- elif finding.priority > 7 -%}\colorbox{ColorRed}{High} -{%- elif finding.priority > 4 -%}\colorbox{ColorOrange}{Medium} -{%- else -%}\colorbox{ColorBlue}{Low} -{%- endif -%}\\ - -\textbf{Issue:}\; {{ issue }}\\ - -\textbf{Recommendation:}\; -{%- if finding.recommendation is not none -%} - {{ recommendation }} -{%- else -%} - Not applicable. -{%- endif -%} - -\begin{figure}[!ht] - \begin{minted}[firstnumber={{ finding.context_start }},highlightlines={{ finding.line_number}}]{python} -{{ finding.context }} - \end{minted} - \caption{\textbf{ {{file}} } on line \textbf{ {{ finding.line_number }} }} -\end{figure} -\newpage -{%- endfor -%} -\end{document} diff --git a/saist/latex/tex/style.tex b/saist/latex/tex/style.tex deleted file mode 100644 index a53430e..0000000 --- a/saist/latex/tex/style.tex +++ /dev/null @@ -1,49 +0,0 @@ -\documentclass[titlepage]{article} -\usepackage[utf8]{inputenc} -\usepackage[margin=0.8in]{geometry} -\usepackage[parfill]{parskip} - -\usepackage[english]{babel} -\usepackage{blindtext} - -\usepackage{fontawesome} -\usepackage[T1]{fontenc} -\renewcommand*\familydefault{\sfdefault} - -\usepackage[table]{xcolor} -\usepackage{booktabs, tabularx} - -\usepackage[cache=false]{minted} -\setminted{ - frame=lines, - autogobble, - linenos, - breaklines, - breakanywhere -} - -\usepackage{graphicx} - -\input{latex/tex/brand.tex} - -\usepackage{titlesec} -\titlespacing{\section} - {0pt}{1\baselineskip}{1.2\baselineskip} - -\titleformat{\section} - {\Large\sffamily\bfseries\color{ColorSecondary}} - {\thesection}{.5em}{} - -\titleformat{\subsection} - {\large\sffamily\color{ColorSecondary}} - {\thesubsection}{.7em}{} - -\usepackage[ - colorlinks=true, - linkcolor=ColorSecondary, % internal document links - urlcolor =ColorPurple, % \href links - citecolor=ColorSecondary, % bibliography -]{hyperref} - -\DeclareUnicodeCharacter{2264}{<=} -\DeclareUnicodeCharacter{2265}{>=} diff --git a/saist/latex/texlive.profile b/saist/latex/texlive.profile deleted file mode 100644 index 32e2c29..0000000 --- a/saist/latex/texlive.profile +++ /dev/null @@ -1,9 +0,0 @@ -selected_scheme scheme-basic -instopt_adjustpath 0 -tlpdbopt_install_docfiles 0 -tlpdbopt_install_srcfiles 0 -TEXDIR /opt/texlive/ -TEXMFLOCAL /opt/texlive/texmf-local -TEXMFSYSCONFIG /opt/texlive/texmf-config -TEXMFSYSVAR /opt/texlive/texmf-var -TEXMFHOME ~/.texmf \ No newline at end of file diff --git a/saist/llm/adapters/__init__.py b/saist/llm/adapters/__init__.py index b6f32ab..786eb8f 100644 --- a/saist/llm/adapters/__init__.py +++ b/saist/llm/adapters/__init__.py @@ -2,35 +2,76 @@ from typing import Callable, List, Optional, Type from pydantic_ai import Agent, Tool - -from typing import Type +from pydantic_ai.settings import ModelSettings from pydantic import BaseModel logger = logging.getLogger("saist.llm.adapters") +THINKING_CHOICES = ("minimal", "low", "medium", "high", "xhigh", "disabled") +DISABLED_THINKING = "disabled" +SAMPLING_PARAMETERS = {"temperature"} + + +def pydantic_ai_supports_thinking() -> bool: + return "thinking" in getattr(ModelSettings, "__annotations__", {}) + + class BaseLlmAdapter: model_options = None model_vendor = '' model_name = '' + def __init__(self, thinking: str = "medium"): + self.thinking = thinking + def get_model_options(self): - return {'temperature': 0.0} | self.model_options if self.model_options is not None else {} + options = {"temperature": 0.0} + if self.model_options is not None: + options.update(self.model_options) + + if pydantic_ai_supports_thinking(): + if self.thinking == DISABLED_THINKING: + options["thinking"] = False + else: + options["thinking"] = self.thinking + for parameter in SAMPLING_PARAMETERS: + options.pop(parameter, None) + elif self.thinking != DISABLED_THINKING: + logger.getChild(self.__class__.__name__).debug( + "Installed pydantic-ai version does not support model_settings.thinking; ignoring setting." + ) + + return options + + async def _run_agent(self, agent: Agent, user_prompt: str): + model_settings = self.get_model_options() + try: + return await agent.run(user_prompt=user_prompt, model_settings=model_settings) + except Exception as e: + if "thinking" not in str(e).lower() or "thinking" not in model_settings: + raise + + fallback_settings = {key: value for key, value in model_settings.items() if key != "thinking"} + logger.getChild(self.__class__.__name__).warning( + "Model or provider rejected model_settings.thinking; retrying without it." + ) + return await agent.run(user_prompt=user_prompt, model_settings=fallback_settings) async def prompt_structured(self, system_prompt: str, user_prompt: str, response_format: Type[BaseModel], tool_fns: Optional[List[Callable]] = None) -> BaseModel: tools = [Tool(fn) for fn in tool_fns] if tool_fns else [] agent = Agent(self.model, output_type = response_format, tools=tools, system_prompt=system_prompt) - response = await agent.run( user_prompt=user_prompt, model_settings=self.get_model_options()) + response = await self._run_agent(agent, user_prompt) logger.getChild(self.__class__.__name__).debug("prompt_structured response", extra={'response_data': response.output, 'prompt': user_prompt}) return response.output async def prompt(self, system_prompt: str, user_prompt: str, tool_fns: Optional[List[Callable]] = None) -> str | None: tools = [Tool(fn) for fn in tool_fns] if tool_fns else [] agent = Agent(self.model, system_prompt=system_prompt, tools=tools) - response = await agent.run(user_prompt=user_prompt, model_settings=self.get_model_options()) + response = await self._run_agent(agent, user_prompt) logger.getChild(self.__class__.__name__).debug("prompt response", extra={'response_data': response.output, 'prompt': user_prompt}) return response.output def generate_agent(self, system_prompt: str = None, tool_fns: Optional[List[Callable]] = None): tools = [Tool(fn) for fn in tool_fns] if tool_fns else [] - return Agent(self.model, system_prompt=system_prompt, tools=tools) + return Agent(self.model, system_prompt=system_prompt, tools=tools, model_settings=self.get_model_options()) diff --git a/saist/llm/adapters/anthropic.py b/saist/llm/adapters/anthropic.py index 20afd78..cd52aa4 100644 --- a/saist/llm/adapters/anthropic.py +++ b/saist/llm/adapters/anthropic.py @@ -6,7 +6,8 @@ from pydantic_ai.providers.anthropic import AnthropicProvider class AnthropicAdapter(BaseLlmAdapter): - def __init__(self, model: str = None, api_key: Optional[str] = None): + def __init__(self, model: str = None, api_key: Optional[str] = None, thinking: str = "medium"): + super().__init__(thinking=thinking) if model is None: model = "claude-3-7-sonnet-latest" self.model = AnthropicModel( diff --git a/saist/llm/adapters/azure_foundry.py b/saist/llm/adapters/azure_foundry.py new file mode 100644 index 0000000..72686f9 --- /dev/null +++ b/saist/llm/adapters/azure_foundry.py @@ -0,0 +1,30 @@ +from typing import Optional + +from llm.adapters import BaseLlmAdapter + +from pydantic_ai.models.openai import OpenAIResponsesModel +from pydantic_ai.providers.azure import AzureProvider + + +class AzureFoundryAdapter(BaseLlmAdapter): + def __init__( + self, + model: str = None, + api_key: Optional[str] = None, + azure_endpoint: Optional[str] = None, + api_version: Optional[str] = None, + thinking: str = "medium", + ): + super().__init__(thinking=thinking) + if model is None: + model = "gpt-5-mini" + self.model = OpenAIResponsesModel( + model, + provider=AzureProvider( + azure_endpoint=azure_endpoint, + api_key=api_key, + api_version=api_version, + ), + ) + self.model_name = self.model.model_name + self.model_vendor = "Azure AI Foundry" diff --git a/saist/llm/adapters/bedrock.py b/saist/llm/adapters/bedrock.py index 006d1c5..35d7603 100644 --- a/saist/llm/adapters/bedrock.py +++ b/saist/llm/adapters/bedrock.py @@ -5,7 +5,8 @@ from pydantic_ai.models.bedrock import BedrockConverseModel class BedrockAdapter(BaseLlmAdapter): - def __init__(self, model: str = None, api_key: Optional[str] = None): + def __init__(self, model: str = None, api_key: Optional[str] = None, thinking: str = "medium"): + super().__init__(thinking=thinking) if api_key: raise ValueError("Do not provide API keys for AWS - use ENV variables") if model is None: @@ -14,4 +15,3 @@ def __init__(self, model: str = None, api_key: Optional[str] = None): self.model_name = self.model.model_name self.model_vendor = 'Bedrock' - diff --git a/saist/llm/adapters/deepseek.py b/saist/llm/adapters/deepseek.py index 211092e..d809404 100644 --- a/saist/llm/adapters/deepseek.py +++ b/saist/llm/adapters/deepseek.py @@ -6,7 +6,8 @@ from pydantic_ai.providers.deepseek import DeepSeekProvider class DeepseekAdapter(BaseLlmAdapter): - def __init__(self, model: str = None, api_key: Optional[str] = None): + def __init__(self, model: str = None, api_key: Optional[str] = None, thinking: str = "medium"): + super().__init__(thinking=thinking) if model is None: model = "deepseek-chat" self.model = OpenAIModel( @@ -15,4 +16,3 @@ def __init__(self, model: str = None, api_key: Optional[str] = None): ) self.model_name = self.model.model_name self.model_vendor = 'DeepSeek' - diff --git a/saist/llm/adapters/faike.py b/saist/llm/adapters/faike.py index c4aa38c..3ee5249 100644 --- a/saist/llm/adapters/faike.py +++ b/saist/llm/adapters/faike.py @@ -10,7 +10,8 @@ logger = logging.getLogger(__name__) class FaikeAdapter(BaseLlmAdapter): - def __init__(self, base_url: str, model: str = None, api_key: Optional[str] = None): + def __init__(self, base_url: str, model: str = None, api_key: Optional[str] = None, thinking: str = "medium"): + super().__init__(thinking=thinking) self.model = model self.model_name = self.model self.model_vendor = 'Fake AI LLM' @@ -18,22 +19,39 @@ def __init__(self, base_url: str, model: str = None, api_key: Optional[str] = No async def prompt_structured(self, system_prompt: str, user_prompt: str, response_format: Type[BaseModel], tool_fns: Optional[List[Callable]] = None) -> BaseModel: logger.getChild(self.__class__.__name__).debug("prompt_structured initial response", extra={'prompt': user_prompt}) - # Extract the real filename from user_prompt, allowing future steps to work - filename: str = re.search("(?<=File: )(.*?)(?=\\n)", user_prompt).group(0) - if response_format is Findings: + # Extract the real filename from user_prompt, allowing future steps to work + filename_match = re.search("(?<=File: )(.*?)(?=\\n)", user_prompt) + if filename_match: + filename = filename_match.group(0) + else: + list_item_match = re.search(r"^- (.+)$", user_prompt, re.MULTILINE) + filename = list_item_match.group(1) if list_item_match else "app.py" fake_finding: dict[str: any] = { "file": filename, "snippet": "[]", "title": "Fake Issue #1234", "issue": "Fake Issue", "recommendation": "Do Nothing", + "validation_steps": ["Inspect the fake finding in app.py.", "Confirm no real validation is required for Faike."], "cwe": "CWE-NAN", "priority": 0, "line_number": 1, } return Findings(findings=[Finding.model_validate(fake_finding)]) + + if hasattr(response_format, "model_fields") and "skills" in response_format.model_fields: + return response_format.model_validate( + { + "skills": [ + { + "filename": "application-routing.md", + "content": "# Application Routing\n\nGenerated by Faike AI for command-path testing.", + } + ] + } + ) else: #Currently prompt_structured only accepts Findings as a return type logger.error(f"Invalid response_format {response_format} passed to prompt_structured (faike.py)") @@ -47,4 +65,3 @@ async def prompt(self, system_prompt: str, user_prompt: str, tool_fns: Optional[ def generate_agent(self, system_prompt: str = None, tool_fns: Optional[List[Callable]] = None): #Faike LLM: Certified non-existent AI doesn't support interactive mode return None - diff --git a/saist/llm/adapters/gemini.py b/saist/llm/adapters/gemini.py index 40de292..a5e097f 100644 --- a/saist/llm/adapters/gemini.py +++ b/saist/llm/adapters/gemini.py @@ -6,7 +6,8 @@ from pydantic_ai.providers.google_gla import GoogleGLAProvider class GeminiAdapter(BaseLlmAdapter): - def __init__(self, model: str = None, api_key: Optional[str] = None): + def __init__(self, model: str = None, api_key: Optional[str] = None, thinking: str = "medium"): + super().__init__(thinking=thinking) if model is None: model = "gemini-2.5-pro-exp-03-25" self.model = GeminiModel( diff --git a/saist/llm/adapters/ollama.py b/saist/llm/adapters/ollama.py index c0e409d..1d5a2f2 100644 --- a/saist/llm/adapters/ollama.py +++ b/saist/llm/adapters/ollama.py @@ -10,7 +10,8 @@ logger = logging.getLogger(__name__) class OllamaAdapter(BaseLlmAdapter): - def __init__(self, base_url: str, model: str = None, api_key: Optional[str] = None): + def __init__(self, base_url: str, model: str = None, api_key: Optional[str] = None, thinking: str = "medium"): + super().__init__(thinking=thinking) if model is None: model = "llama3:latest" self.model = model @@ -68,4 +69,3 @@ async def prompt(self, system_prompt: str, user_prompt: str, tool_fns: Optional[ def generate_agent(self, system_prompt: str = None, tool_fns: Optional[List[Callable]] = None): #TODO: OLLAMA AGENT FOR SHELL return None - diff --git a/saist/llm/adapters/openai.py b/saist/llm/adapters/openai.py index 6bca158..b992513 100644 --- a/saist/llm/adapters/openai.py +++ b/saist/llm/adapters/openai.py @@ -2,16 +2,23 @@ from llm.adapters import BaseLlmAdapter -from pydantic_ai.models.openai import OpenAIModel +from pydantic_ai.models.openai import OpenAIResponsesModel from pydantic_ai.providers.openai import OpenAIProvider class OpenAiAdapter(BaseLlmAdapter): - def __init__(self, model: str = None, api_key: Optional[str] = None): + def __init__( + self, + model: str = None, + api_key: Optional[str] = None, + base_url: Optional[str] = None, + thinking: str = "medium", + ): + super().__init__(thinking=thinking) if model is None: model = "gpt-4o" - self.model = OpenAIModel( + self.model = OpenAIResponsesModel( model, - provider = OpenAIProvider( api_key=api_key ) + provider=OpenAIProvider(api_key=api_key, base_url=base_url), ) self.model_name = self.model.model_name self.model_vendor = 'OpenAI' diff --git a/saist/main.py b/saist/main.py index 114881f..6d425e2 100755 --- a/saist/main.py +++ b/saist/main.py @@ -2,13 +2,14 @@ import asyncio import logging import os -from typing import Optional +from typing import Callable, Optional from dotenv import load_dotenv -from latex import Latex +from reportlab_pdf import ReportLabPdf from llm.adapters import BaseLlmAdapter from llm.adapters.anthropic import AnthropicAdapter +from llm.adapters.azure_foundry import AzureFoundryAdapter from llm.adapters.bedrock import BedrockAdapter from llm.adapters.deepseek import DeepseekAdapter from llm.adapters.faike import FaikeAdapter @@ -32,6 +33,14 @@ from util.output import print_banner, write_csv from util.poem import poem from util.prompts import prompts +from util.skills import ( + format_analysis_skills, + generate_skill_files, + load_analysis_skills, + project_root_from_args, + resolve_skills_dir, + skills_prompt_digest, +) from web import FindingsServer @@ -44,17 +53,59 @@ logger = logging.getLogger("saist") -async def analyze_single_file(scm: Scm, adapter: BaseLlmAdapter, filename, patch_text, disable_tools: bool) -> Optional[list[Finding]]: + +class CoverageTrackingScm: + def __init__(self, scm: Scm): + self.scm = scm + self.files_read: set[str] = set() + + async def read_file_contents(self, filename: str): + contents = await self.scm.read_file_contents(filename) + if contents is not None: + self.files_read.add(filename) + return contents + + async def list_files(self) -> list[str]: + return await self.scm.list_files() + + async def regex_search( + self, + pattern: str, + file_pattern: str = "**/*", + max_results: int = 100, + ) -> list[dict[str, str | int]]: + results = await self.scm.regex_search(pattern, file_pattern, max_results) + for result in results: + filename = result.get("filename") if isinstance(result, dict) else None + if filename: + self.files_read.add(str(filename)) + return results + + def tool_functions(self) -> list[Callable]: + return [self.read_file_contents, self.list_files, self.regex_search] + + +def scm_detect_prompt(scm: Scm) -> str: + return scm.detect_prompt() if hasattr(scm, "detect_prompt") else FilesystemAdapter.DETECT_PROMPT + + +def scm_summary_prompt(scm: Scm) -> str: + return scm.summary_prompt() if hasattr(scm, "summary_prompt") else FilesystemAdapter.SUMMARY_PROMPT + +async def analyze_single_file(scm: Scm, adapter: BaseLlmAdapter, filename, patch_text, disable_tools: bool, analysis_skills: str = "") -> Optional[list[Finding]]: """ Analyzes a SINGLE file diff with OpenAI, returning a Findings object or None on error. """ - system_prompt = prompts.DETECT + system_prompt = prompts.detect(scm_detect_prompt(scm)) + if analysis_skills: + system_prompt = f"{system_prompt}\n\n{analysis_skills}" + logger.debug(f"Processing {filename}") prompt = ( f"\n\nFile: {filename}\n{patch_text}\n" ) try: - return (await adapter.prompt_structured(system_prompt, prompt, Findings, [] if disable_tools else [scm.read_file_contents])).findings + return (await adapter.prompt_structured(system_prompt, prompt, Findings, [] if disable_tools else scm.tool_functions())).findings except Exception as e: logger.error(f"[Error] File '{filename}': {e}") return None @@ -79,14 +130,20 @@ async def context_from_finding(scm: Scm, finding: Finding, context_size: int = 3 return "\n".join(context), start, end -def generate_summary_from_findings(adapter: BaseLlmAdapter, findings: list[Finding]) -> str: +def generate_summary_from_findings(adapter: BaseLlmAdapter, findings: list[Finding], scm_prompt: str = "") -> str: """ Uses OpenAI to generate a summary of all findings to be used as the PR review body. """ - system_prompt = prompts.SUMMARY + system_prompt = prompts.summary(scm_prompt) prompt = "" for f in findings: - prompt += f"- **File**: `{f.file}`\n - **Issue**: {f.issue}\n - **Recommendation**: {f.recommendation}\n\n" + validation_steps = "\n".join(f" - {step}" for step in f.validation_steps) or " - Not provided" + prompt += ( + f"- **File**: `{f.file}`\n" + f" - **Issue**: {f.issue}\n" + f" - **Recommendation**: {f.recommendation}\n" + f" - **Validation steps**:\n{validation_steps}\n\n" + ) try: return adapter.prompt(system_prompt, prompt) @@ -94,6 +151,211 @@ def generate_summary_from_findings(adapter: BaseLlmAdapter, findings: list[Findi logger.error(f"[Error generating summary] {e}") return "Security issues found. Please review the inline comments." + +def build_finding_review_body(finding: Finding) -> str: + priority = "LOW" + if finding.priority > 4: + priority = "MEDIUM" + if finding.priority > 7: + priority = "HIGH" + if finding.priority > 8: + priority = "CRITICAL" + + return ( + f"**Security Issue:** {finding.issue}\n\n" + f"**Priority:** {priority}\n\n" + f"**CWE:** {finding.cwe}\n\n" + f"**Recommendation:** {finding.recommendation or 'None provided.'}\n\n" + f"**Validation Steps:**\n{format_validation_steps(finding.validation_steps)}\n\n" + f"**Snippet**: `{finding.snippet}`\n\n" + ) + + +def format_validation_steps(validation_steps: list[str]) -> str: + if not validation_steps: + return "None provided." + return "\n".join(f"{index}. {step}" for index, step in enumerate(validation_steps, start=1)) + + +def build_filesystem_review_comments(findings: list[Finding]) -> list[dict]: + comments = [] + for finding in findings: + if not finding.file or not finding.snippet or not finding.issue: + continue + comments.append( + { + "path": finding.file, + "position": max(1, finding.line_number), + "body": build_finding_review_body(finding), + } + ) + return comments + + +def build_diff_review_comments(findings: list[Finding], file_line_maps: dict) -> list[dict]: + comments = [] + for finding in findings: + diff_position = file_line_maps[finding.file][finding.line_number] + comments.append( + { + "path": finding.file, + "position": diff_position - 1, + "body": build_finding_review_body(finding), + } + ) + return comments + + +def dedupe_findings(findings: list[Finding]) -> list[Finding]: + deduped: dict[tuple[str, int], Finding] = {} + order: list[tuple[str, int]] = [] + + for finding in findings: + key = (finding.file, finding.line_number) + existing = deduped.get(key) + if existing is None: + deduped[key] = finding + order.append(key) + continue + + if finding.priority > existing.priority: + deduped[key] = finding + + return [deduped[key] for key in order] + + +async def generate_findings_with_filesystem_tools( + scm: Scm, + llm: BaseLlmAdapter, + filenames: list[str], + disable_tools: bool, + analysis_skills: str = "", +) -> tuple[list[Finding], set[str]]: + tracked_scm = CoverageTrackingScm(scm) + system_prompt = prompts.detect(scm_detect_prompt(scm)) + if analysis_skills: + system_prompt = f"{system_prompt}\n\n{analysis_skills}" + + file_list = "\n".join(f"- {filename}" for filename in filenames) + user_prompt = f""" +Application file inventory: +{file_list} + +Perform a penetration test style application security review of this codebase. +Use the available tools to inspect files before reporting findings. +Return only findings that are supported by code you inspected. +""" + + result = await llm.prompt_structured( + system_prompt, + user_prompt, + Findings, + [] if disable_tools else tracked_scm.tool_functions(), + ) + return result.findings, tracked_scm.files_read + + +async def generate_findings_with_filesystem_tools_iterations( + scm: Scm, + llm: BaseLlmAdapter, + filenames: list[str], + disable_tools: bool, + analysis_skills: str, + iterations: int, + max_concurrent: int, + disable_caching: bool = True, + cache_folder: str = ".cache", +) -> tuple[list[Finding], set[str]]: + semaphore = asyncio.Semaphore(max(1, max_concurrent)) + if disable_caching is False: + os.makedirs(cache_folder, exist_ok=True) + + overall_progress = Progress( + TextColumn("[bold blue]{task.description}"), + BarColumn(), + MofNCompleteColumn(), + TimeElapsedColumn(), + TimeRemainingColumn(), + ) + + iteration_progress = Progress( + SpinnerColumn(), + TextColumn("[blue]{task.description}"), + transient=True, + ) + + progress_group = Group( + overall_progress, + iteration_progress, + ) + + async def run_uncached_iteration(iteration: int): + return await generate_findings_with_filesystem_tools( + scm=scm, + llm=llm, + filenames=filenames, + disable_tools=disable_tools, + analysis_skills=analysis_skills, + ) + + async def run_cached_iteration(iteration: int): + cache_hash = await hash_files(scm, filenames, extra=analysis_skills) + cache_file = os.path.join(cache_folder, f"{iteration}-{cache_hash}.json") + if os.path.exists(cache_file): + return filesystem_tool_findings_from_cache_file(cache_file) + + findings, files_read = await run_uncached_iteration(iteration) + store_filesystem_tool_findings_to_cache_file( + iteration=iteration, + filenames=filenames, + findings=findings, + files_read=files_read, + cache_file=cache_file, + ) + return findings, files_read + + async def run_iteration(iteration: int, overall_task): + async with semaphore: + iteration_task = iteration_progress.add_task( + description=f"Iteration {iteration}/{iterations}...", + transient=True, + ) + try: + if disable_caching: + return await run_uncached_iteration(iteration) + return await run_cached_iteration(iteration) + finally: + iteration_progress.remove_task(iteration_task) + iteration_progress.refresh() + overall_progress.update(overall_task, advance=1) + + with Live(progress_group): + overall_task = overall_progress.add_task( + f"Running {iterations} tool-driven analysis iteration{'s' if iterations != 1 else ''}...", + total=iterations, + start=True, + ) + try: + results = await asyncio.gather(*(run_iteration(iteration, overall_task) for iteration in range(1, iterations + 1))) + finally: + overall_progress.stop() + iteration_progress.stop() + + all_findings = [] + files_read = set() + for findings, iteration_files_read in results: + all_findings.extend(findings) + files_read.update(iteration_files_read) + + return all_findings, files_read + + +def print_coverage(files_read: set[str], files_in_scope: list[str]): + total = len(files_in_scope) + read_count = len(files_read.intersection(files_in_scope)) + percent = (read_count / total * 100) if total else 0 + print(f"๐Ÿ“ˆ LLM file coverage: {read_count}/{total} files read ({percent:.1f}%)\n") + def _get_scm_adapter(args) -> BaseScmAdapter: if args.SCM == 'github': logger.debug("Using SCM: Github") @@ -122,28 +384,43 @@ def _get_scm_adapter(args) -> BaseScmAdapter: async def _get_llm_adapter(args) -> BaseLlmAdapter: model = args.llm_model + thinking = args.thinking if args.llm == 'anthropic': - llm = AnthropicAdapter( api_key = args.llm_api_key, model=model) + llm = AnthropicAdapter(api_key=args.llm_api_key, model=model, thinking=thinking) logger.debug(f"Using LLM: anthropic Model: {llm.model_name}") + elif args.llm == 'azure-foundry': + llm = AzureFoundryAdapter( + api_key=args.llm_api_key, + model=model, + azure_endpoint=args.azure_openai_endpoint, + api_version=args.azure_openai_api_version, + thinking=thinking, + ) + logger.debug(f"Using LLM: Azure AI Foundry Model: {llm.model_name}") elif args.llm == 'bedrock': - llm = BedrockAdapter( api_key = args.llm_api_key, model=model) + llm = BedrockAdapter(api_key=args.llm_api_key, model=model, thinking=thinking) logger.debug(f"Using LLM: AWS bedrock Model: {llm.model_name}") elif args.llm == 'deepseek': - llm = DeepseekAdapter(api_key = args.llm_api_key, model=model) + llm = DeepseekAdapter(api_key=args.llm_api_key, model=model, thinking=thinking) logger.debug(f"Using LLM: deepseek Model: {llm.model_name}") elif args.llm == 'openai': - llm = OpenAiAdapter(api_key = args.llm_api_key, model=model) + llm = OpenAiAdapter( + api_key=args.llm_api_key, + model=model, + base_url=args.openai_base_uri, + thinking=thinking, + ) logger.debug(f"Using LLM: openai Model: {llm.model_name}") elif args.llm == 'gemini': - llm = GeminiAdapter(api_key = args.llm_api_key, model=model) + llm = GeminiAdapter(api_key=args.llm_api_key, model=model, thinking=thinking) logger.debug(f"Using LLM: gemini Model: {llm.model_name}") elif args.llm == 'ollama': - llm = OllamaAdapter(api_key = args.llm_api_key, base_url=args.ollama_base_uri, model=model) + llm = OllamaAdapter(api_key=args.llm_api_key, base_url=args.ollama_base_uri, model=model, thinking=thinking) await llm.initialize() logger.debug(f"Using LLM: ollama Model: {llm.model_name}") elif args.llm == 'faike': - llm = FaikeAdapter("", "Fake LLM") + llm = FaikeAdapter("", "Fake LLM", thinking=thinking) logger.debug("Using LLM: Faike AI") else: raise Exception("Could not determine a suitable LLM to use") @@ -179,140 +456,208 @@ async def main(): print("โœจ Poem generation completed.\n") exit(0) + project_root = project_root_from_args(args) + skills_dir = resolve_skills_dir(project_root, args.skills_path) + + if args.generate_skills: + print("๐Ÿงญ Generating SAIST analysis skills...") + result = await generate_skill_files( + llm=llm, + project_root=project_root, + skills_dir=skills_dir, + max_files=args.skills_sample_files, + max_file_bytes=args.skills_sample_bytes, + overwrite=args.overwrite_skills, + ) + print(f"โœ… Sampled {result.sampled_files} files from {project_root}") + print(f"โœ… Wrote {len(result.written)} skill files to {result.skills_dir}") + if result.skipped: + print(f"โ„น๏ธ Skipped {len(result.skipped)} existing skill files. Use --overwrite-skills to replace them.") + return + print("๐Ÿ”Ž Initializing SCM adapter...") scm_adapter = _get_scm_adapter(args) scm = Scm(adapter=scm_adapter) print(f"โœ… Using SCM: {args.SCM}\n") - # 1) Get changed files - print("๐Ÿ“‚ Fetching changed files...") - changed_files = scm.get_changed_files() - if not changed_files: - print("โš ๏ธ No changed files detected. Exiting.") - return - print(f"โœ… Detected {len(changed_files)} changed files\n") - - # 2) Gather only relevant app code diffs - print("๐Ÿงน Filtering relevant app code diffs...") - file_line_maps = {} - file_new_lines_text = {} - app_files = [] + shallow_filesystem_scan = args.SCM == "filesystem" and not args.deep + analysis_skills = "" + if not args.disable_skills: + loaded_skills = load_analysis_skills(skills_dir, args.skills_max_bytes) + if shallow_filesystem_scan and not loaded_skills: + print("๐Ÿงญ No SAIST skill files found. Generating application skills for this filesystem scan...") + result = await generate_skill_files( + llm=llm, + project_root=project_root, + skills_dir=skills_dir, + max_files=args.skills_sample_files, + max_file_bytes=args.skills_sample_bytes, + overwrite=False, + ) + print(f"โœ… Sampled {result.sampled_files} files from {project_root}") + print(f"โœ… Wrote {len(result.written)} skill files to {result.skills_dir}") + if result.skipped: + print(f"โ„น๏ธ Skipped {len(result.skipped)} existing skill files.") + loaded_skills = load_analysis_skills(skills_dir, args.skills_max_bytes) + + if loaded_skills: + analysis_skills = format_analysis_skills(loaded_skills) + print(f"๐Ÿง  Loaded {len(loaded_skills)} SAIST skill files from {skills_dir}\n") + else: + logging.debug(f"No SAIST skill files found under {skills_dir}") + else: + logging.debug("SAIST skill loading disabled") filter_rules = FilterRules(args.include, args.exclude) - for f in changed_files: - filename = f["filename"] - patch_text = f.get("patch", "") - if not patch_text: - logging.debug(f"Skipped file {filename} as it contains no patch text") - continue - - if not filter_rules.filename_included(filename): - logging.debug(f"Skipped file {filename} as it is not included in rules") - continue - - if not args.skip_line_length_check: - if filter_rules.file_exceeds_line_length_limit(file_content=await scm.read_file_contents(filename), patch_text=patch_text, max_line_length=args.max_line_length): - logging.debug(f"Skipped file {filename} as it contains lines that exceed the maximum line length ({args.max_line_length})") - continue - - line_map, new_lines_text = parse_unified_diff(patch_text) - file_line_maps[filename] = line_map - file_new_lines_text[filename] = new_lines_text - app_files.append((filename, patch_text)) - - if not app_files: - print("โš ๏ธ No app code diffs to analyze. Exiting.") - return - print(f"โœ… Prepared {len(app_files)} app files for analysis.\n") + review_comments = [] - app_filenames = list((filename for filename,_ in app_files)) - logging.debug(f"Files to process: {app_filenames}") + if shallow_filesystem_scan: + print("๐Ÿ“‚ Listing application files...") + all_files = await scm.list_files() + app_filenames = [filename for filename in all_files if filter_rules.filename_included(filename)] + + if not app_filenames: + print("โš ๏ธ No app files to analyze. Exiting.") + return + + print(f"โœ… Prepared {len(app_filenames)} app files for tool-driven analysis.\n") + + if args.dry_run: + print("โš ๏ธ --dry-run flag passed, exiting without analyzing files.") + exit(0) + + print(f"๐Ÿ” Analyzing application with LLM tool use ({args.iterations} iteration{'s' if args.iterations != 1 else ''})...") + all_findings, files_read = await generate_findings_with_filesystem_tools_iterations( + scm=scm, + llm=llm, + filenames=app_filenames, + disable_tools=args.disable_tools, + analysis_skills=analysis_skills, + iterations=args.iterations, + max_concurrent=args.llm_rate_limit, + disable_caching=args.disable_caching, + cache_folder=args.cache_folder, + ) + print_coverage(files_read, app_filenames) - if args.dry_run: - print("โš ๏ธ --dry-run flag passed, exiting without analyzing files.") - exit(0) + if not all_findings: + print("โœ… No findings reported. Exiting.\n") + return - # 3) Analyze each file in parallel - print("๐Ÿ” Analyzing files for security issues...") - max_workers = min(args.llm_rate_limit, len(app_files)) - logging.debug(f"{max_workers=}") - all_findings = await generate_findings(scm, llm, app_files, max_workers, args.disable_tools, args.disable_caching, args.cache_folder) + else: - if not all_findings: - print("โœ… No findings reported. Exiting.\n") - return - print(f"๐Ÿšจ Analysis complete! Identified {len(all_findings)} potential issues.\n") + # 1) Get changed files + print("๐Ÿ“‚ Fetching changed files...") + changed_files = scm.get_changed_files() + if not changed_files: + print("โš ๏ธ No changed files detected. Exiting.") + return + print(f"โœ… Detected {len(changed_files)} changed files\n") + + # 2) Gather only relevant app code diffs + print("๐Ÿงน Filtering relevant app code diffs...") + file_line_maps = {} + file_new_lines_text = {} + app_files = [] + + for f in changed_files: + filename = f["filename"] + patch_text = f.get("patch", "") + if not patch_text: + logging.debug(f"Skipped file {filename} as it contains no patch text") + continue + + if not filter_rules.filename_included(filename): + logging.debug(f"Skipped file {filename} as it is not included in rules") + continue - # 4) Build review comments from snippet-based findings - review_comments = [] - all_findings.sort(key=lambda x: x.priority,reverse=True) - - for item in all_findings: - item.line_number = -1 #set to -1 for filtering. Gets changed later if finding is valid - file_name = item.file - snippet = item.snippet - issue = item.issue - priority = "LOW" - if item.priority > 4: - priority = "MEDIUM" - if item.priority > 7: - priority = "HIGH" - if item.priority > 8: - priority = "CRITICAL" - cwe = item.cwe - recommendation = item.recommendation - - # Basic checks - if not file_name or not snippet or not issue: - continue - if file_name not in file_line_maps: - # Possibly flagged a file that doesn't exist in the PR - continue + if not args.skip_line_length_check: + if filter_rules.file_exceeds_line_length_limit(file_content=await scm.read_file_contents(filename), patch_text=patch_text, max_line_length=args.max_line_length): + logging.debug(f"Skipped file {filename} as it contains lines that exceed the maximum line length ({args.max_line_length})") + continue + + line_map, new_lines_text = parse_unified_diff(patch_text) + file_line_maps[filename] = line_map + file_new_lines_text[filename] = new_lines_text + app_files.append((filename, patch_text)) + + if not app_files: + print("โš ๏ธ No app code diffs to analyze. Exiting.") + return + print(f"โœ… Prepared {len(app_files)} app files for analysis.\n") + + app_filenames = list((filename for filename,_ in app_files)) + logging.debug(f"Files to process: {app_filenames}") + + if args.dry_run: + print("โš ๏ธ --dry-run flag passed, exiting without analyzing files.") + exit(0) + + # 3) Analyze each file in parallel + print("๐Ÿ” Analyzing files for security issues...") + max_workers = min(args.llm_rate_limit, len(app_files)) + logging.debug(f"{max_workers=}") + all_findings = await generate_findings(scm, llm, app_files, max_workers, args.disable_tools, args.disable_caching, args.cache_folder, analysis_skills) + + if not all_findings: + print("โœ… No findings reported. Exiting.\n") + return + + # 4) Build review comments from snippet-based findings + all_findings.sort(key=lambda x: x.priority,reverse=True) + + for item in all_findings: + item.line_number = -1 #set to -1 for filtering. Gets changed later if finding is valid + file_name = item.file + snippet = item.snippet + issue = item.issue + + # Basic checks + if not file_name or not snippet or not issue: + continue + if file_name not in file_line_maps: + # Possibly flagged a file that doesn't exist in the PR + continue - line_map = file_line_maps[file_name] - new_lines_text = file_new_lines_text[file_name] + new_lines_text = file_new_lines_text[file_name] - # Attempt to find which 'new_line' has the snippet - matched_new_line = None - for ln, code_text in new_lines_text.items(): - if snippet in code_text: - matched_new_line = ln - break + # Attempt to find which 'new_line' has the snippet + matched_new_line = None + for ln, code_text in new_lines_text.items(): + if snippet in code_text: + matched_new_line = ln + break - if not matched_new_line: - # If we can't find the snippet in the patch, skip - continue + if not matched_new_line: + # If we can't find the snippet in the patch, skip + continue - diff_position = line_map[matched_new_line] - body_text = ( - f"**Security Issue:** {issue}\n\n" - f"**Priority:** {priority}\n\n" - f"**CWE:** {cwe}\n\n" - f"**Recommendation:** {recommendation or 'None provided.'}\n\n" - f"**Snippet**: `{snippet}`\n\n" - ) + item.line_number = matched_new_line - review_comments.append({ - "path": file_name, - "position": diff_position - 1, - "body": body_text - }) + all_findings = list([x for x in all_findings if x.line_number != -1]) - item.line_number = matched_new_line - all_findings = list([x for x in all_findings if x.line_number != -1]) + all_findings = dedupe_findings(all_findings) + all_findings.sort(key=lambda x: x.priority, reverse=True) if not all_findings: print("No issues detected") exit(0) + print(f"๐Ÿšจ Analysis complete! Identified {len(all_findings)} potential issues.\n") + + if shallow_filesystem_scan: + review_comments = build_filesystem_review_comments(all_findings) + else: + review_comments = build_diff_review_comments(all_findings, file_line_maps) + if args.interactive: s = Shell(llm, scm, all_findings) await s.run() all_findings = s.findings - comment = await generate_summary_from_findings(llm, all_findings) + comment = await generate_summary_from_findings(llm, all_findings, scm_summary_prompt(scm)) scm.create_review( comment=comment, review_comments=review_comments, @@ -345,7 +690,7 @@ async def main(): w = FindingsServer(args.web_host, args.web_port) w.run(enriched_findings) - if args.pdf or args.tex: + if args.pdf: findings_context = [] for finding in all_findings: try: @@ -361,21 +706,25 @@ async def main(): findings_context.append(fc) except: continue - l = Latex(llm, args.project_name, findings_context, comment) - l.run(args) + r = ReportLabPdf(llm, args.project_name, findings_context, comment) + r.run(args) if args.ci and len(all_findings) > 0: exit(1) -async def process_file(scm: Scm, llm, filename, patch_text, disable_tools, disable_caching, cache_folder): +async def process_file(scm: Scm, llm, filename, patch_text, disable_tools, disable_caching, cache_folder, analysis_skills): start = asyncio.get_event_loop().time() if disable_caching is True: - result = await analyze_single_file(scm, llm, filename, patch_text, disable_tools) + result = await analyze_single_file(scm, llm, filename, patch_text, disable_tools, analysis_skills) else: hash: str = await hash_file(scm, filename) - cache_file = os.path.join(cache_folder, hash + ".json") + cache_filename = hash + ".json" + if analysis_skills: + cache_filename = hash + "-" + skills_prompt_digest(analysis_skills) + ".json" + + cache_file = os.path.join(cache_folder, cache_filename) if not os.path.exists(cache_file): - result = await analyze_single_file(scm, llm, filename, patch_text, disable_tools) + result = await analyze_single_file(scm, llm, filename, patch_text, disable_tools, analysis_skills) store_findings_to_cache_file(filename, result, cache_file) else: result = findings_from_cache_file(cache_file) @@ -385,7 +734,7 @@ async def process_file(scm: Scm, llm, filename, patch_text, disable_tools, disab return result -async def generate_findings(scm, llm, app_files, max_concurrent, disable_tools, disable_caching, cache_folder): +async def generate_findings(scm, llm, app_files, max_concurrent, disable_tools, disable_caching, cache_folder, analysis_skills): if disable_caching is False: if not os.path.exists(cache_folder) or not os.path.isdir(cache_folder): os.makedirs(cache_folder, exist_ok=True) @@ -427,7 +776,7 @@ async def sub_func(*args, **kwargs): tasks = [] overall_task = overall_progress.add_task(f"Analyzing {len(app_files)} files...", total=len(app_files), start=True) # Add a task for filename, patch_text in app_files: - wrapper_func = task_progress_wrapper(process_file, overall_progress, overall_task, file_progress, filename, semaphore)(scm, llm, filename, patch_text, disable_tools, disable_caching, cache_folder) + wrapper_func = task_progress_wrapper(process_file, overall_progress, overall_task, file_progress, filename, semaphore)(scm, llm, filename, patch_text, disable_tools, disable_caching, cache_folder, analysis_skills) tasks.append( wrapper_func ) diff --git a/saist/models.py b/saist/models.py index 0c80b8d..74f7706 100644 --- a/saist/models.py +++ b/saist/models.py @@ -8,6 +8,13 @@ class Finding(BaseModel): title: Annotated[str, Field(description= "a short title describing the issue") ] issue: str recommendation: str + validation_steps: Annotated[ + list[str], + Field( + default_factory=list, + description="concrete steps a reviewer can follow to validate the issue is exploitable", + ), + ] cwe: Annotated[str, Field(description= "CWE id, should conform to CWE-XX or CWE-XXX where X is a number") ] priority: int line_number: int @@ -27,4 +34,4 @@ class FindingEnriched(Finding): class FindingContext(Finding): context: str context_start: int - context_end: int \ No newline at end of file + context_end: int diff --git a/saist/reportlab_pdf/__init__.py b/saist/reportlab_pdf/__init__.py new file mode 100644 index 0000000..b260497 --- /dev/null +++ b/saist/reportlab_pdf/__init__.py @@ -0,0 +1,789 @@ +from dataclasses import dataclass +from html import escape +import datetime +import logging +import os +import textwrap + +from llm.adapters import BaseLlmAdapter +from models import FindingContext +from reportlab.lib import colors +from reportlab.lib.enums import TA_CENTER +from reportlab.lib.pagesizes import A4 +from reportlab.lib.styles import ParagraphStyle, getSampleStyleSheet +from reportlab.lib.units import inch +from reportlab.pdfbase import pdfmetrics +from reportlab.pdfbase.ttfonts import TTFont +from reportlab.platypus import Image, PageBreak, Paragraph, SimpleDocTemplate, Spacer, Table, TableStyle +from reportlab.platypus.tableofcontents import TableOfContents + +logger = logging.getLogger("saist.reportlab_pdf") + +PUNK_BG_HEX = "#111B29" +PUNK_SECONDARY_HEX = "#0C2540" +PUNK_SECONDARY_LIGHT_HEX = "#2E3848" +PUNK_TEXT_HEX = "#DDEEF2" +PUNK_MUTED_HEX = "#AFC3CB" +PUNK_BLUE_HEX = "#5BBFCF" +PUNK_PURPLE_HEX = "#A25DCE" +PUNK_ORANGE_HEX = "#CE985D" +PUNK_GREEN_HEX = "#87CE5D" +PUNK_RED_HEX = "#CE5D5D" +PUNK_BG = colors.HexColor(PUNK_BG_HEX) +PUNK_SECONDARY = colors.HexColor(PUNK_SECONDARY_HEX) +PUNK_SECONDARY_LIGHT = colors.HexColor(PUNK_SECONDARY_LIGHT_HEX) +PUNK_TEXT = colors.HexColor(PUNK_TEXT_HEX) +PUNK_MUTED = colors.HexColor(PUNK_MUTED_HEX) +PUNK_BLUE = colors.HexColor(PUNK_BLUE_HEX) +PUNK_PURPLE = colors.HexColor(PUNK_PURPLE_HEX) +PUNK_ORANGE = colors.HexColor(PUNK_ORANGE_HEX) +PUNK_GREEN = colors.HexColor(PUNK_GREEN_HEX) +PUNK_RED = colors.HexColor(PUNK_RED_HEX) +PUNK_WHITE = colors.white +PRINT_BG = colors.white +PRINT_TEXT = colors.HexColor("#132033") +PRINT_MUTED = colors.HexColor("#52616F") +PRINT_PANEL = colors.HexColor("#F5F8FA") +PRINT_PANEL_ALT = colors.HexColor("#EEF4F7") +PRINT_BORDER = colors.HexColor("#D8E3E8") +PRINT_CODE_BG = colors.HexColor("#F7FAFC") +PRINT_CODE_HIGHLIGHT = colors.HexColor("#E8F6E8") +RAINBOW_HEX = [PUNK_BLUE_HEX, PUNK_PURPLE_HEX, PUNK_ORANGE_HEX, PUNK_GREEN_HEX, PUNK_RED_HEX] +RAINBOW = [colors.HexColor(color) for color in RAINBOW_HEX] +CONTENT_WIDTH = 6.7 * inch +ISSUE_LABEL_WIDTH = 1.0 * inch +CODE_LINE_NUMBER_WIDTH = 0.62 * inch +CODE_WRAP_CHARS = 112 +FRAME_INNER_WIDTH = A4[0] - (2 * 0.65 * inch) - 12 +CONTENT_RIGHT_INDENT = max(0, FRAME_INNER_WIDTH - CONTENT_WIDTH) + + +def _font_asset_path(filename: str) -> str: + return os.path.abspath(os.path.join(os.path.dirname(__file__), "fonts", filename)) + + +def _register_pathway_extreme() -> str: + font_path = _font_asset_path("PathwayExtreme.ttf") + if not os.path.exists(font_path): + logger.warning("Pathway Extreme font file missing, falling back to Helvetica") + return "Helvetica" + + try: + pdfmetrics.registerFont(TTFont("PathwayExtreme", font_path)) + return "PathwayExtreme" + except Exception as e: + logger.warning(f"Could not register Pathway Extreme font, falling back to Helvetica: {e}") + return "Helvetica" + + +FONT_REGULAR = _register_pathway_extreme() +FONT_BOLD = FONT_REGULAR if FONT_REGULAR != "Helvetica" else "Helvetica-Bold" +FONT_MONO = "Courier" + + +class SaistDocTemplate(SimpleDocTemplate): + def afterFlowable(self, flowable): + level = getattr(flowable, "_saist_toc_level", None) + if level is not None: + key = getattr(flowable, "_saist_toc_key", None) + if key: + self.canv.bookmarkPage(key) + self.notify("TOCEntry", (level, flowable.getPlainText(), self.page, key)) + + +@dataclass +class ReportLabPdf: + _DEFAULT_OUTPUT_DIR = "reporting" + + llm: BaseLlmAdapter + project: str + findings: list[FindingContext] + comment: str + + def run(self, args): + pdf_path = os.path.join(self._DEFAULT_OUTPUT_DIR, args.pdf_filename) + self.generated_at = datetime.datetime.now().isoformat(sep=" ", timespec="minutes") + self.scm_adapter = self._scm_adapter_from_args(args) + self.target = self._target_from_args(args) + + try: + os.makedirs(os.path.dirname(pdf_path), exist_ok=True) + doc = SaistDocTemplate( + pdf_path, + pagesize=A4, + rightMargin=0.65 * inch, + leftMargin=0.65 * inch, + topMargin=0.65 * inch, + bottomMargin=0.65 * inch, + title=self._document_title(), + author="SAIST", + ) + doc.multiBuild(self._story(), onFirstPage=self._draw_page, onLaterPages=self._draw_page) + except Exception as e: + logger.error(f"Unable to write ReportLab PDF file to '{pdf_path}': {e}") + exit(1) + + print(f"Written ReportLab PDF report to '{pdf_path}'") + + def _story(self) -> list: + styles = self._styles() + story = self._cover_story(styles) + + story.extend( + [ + PageBreak(), + Paragraph("Report contents", styles["Heading1"]), + self._table_of_contents(styles), + PageBreak(), + self._toc_heading("Summary", styles["Heading1"], 0, "summary"), + self._text_panel(self.comment or "No summary was generated.", styles), + Spacer(1, 0.28 * inch), + Paragraph("Issue summary", styles["Heading2"]), + self._issue_summary_table(styles), + PageBreak(), + self._toc_heading("Issues", styles["Heading1"], 0, "issues"), + Paragraph( + self._escaped(f"{len(self.findings)} issues reviewed with code context and remediation guidance."), + styles["Body"], + ), + Spacer(1, 0.18 * inch), + ] + ) + + if not self.findings: + story.append(self._text_panel("No issues were provided.", styles)) + return story + + for index, finding in enumerate(self.findings, start=1): + if index > 1: + story.append(PageBreak()) + + story.extend(self._finding_story(index, finding, styles)) + + return story + + def _cover_story(self, styles: dict[str, ParagraphStyle]) -> list: + story = [] + logo_path = self._logo_path() + if logo_path: + logo = Image(logo_path, width=1.15 * inch, height=1.40 * inch) + logo.hAlign = "CENTER" + story.extend([Spacer(1, 0.25 * inch), logo, Spacer(1, 0.20 * inch)]) + else: + story.append(Spacer(1, 0.70 * inch)) + + story.extend( + [ + Paragraph("Punk Security", styles["Title"]), + Spacer(1, 0.12 * inch), + Paragraph(self._cover_title_markup(), styles["CoverSubtitle"]), + Spacer(1, 0.30 * inch), + self._metadata_table(styles), + Spacer(1, 0.18 * inch), + self._disclaimer_panel(styles), + ] + ) + + return story + + def _table_of_contents(self, styles: dict[str, ParagraphStyle]) -> TableOfContents: + toc = TableOfContents() + toc.levelStyles = [styles["TocLevel0"], styles["TocLevel1"]] + toc.dotsMinLevel = 0 + return toc + + def _finding_story(self, index: int, finding: FindingContext, styles: dict[str, ParagraphStyle]) -> list: + priority_label, priority_color = self._priority(finding.priority) + metadata = [ + ["Priority", priority_label], + ["CWE", finding.cwe or "Not specified"], + ["File", finding.file], + ["Line", str(finding.line_number)], + ] + + table = Table(metadata, colWidths=[ISSUE_LABEL_WIDTH, CONTENT_WIDTH - ISSUE_LABEL_WIDTH], hAlign="LEFT") + table.setStyle( + TableStyle( + [ + ("BACKGROUND", (0, 0), (0, -1), PRINT_PANEL_ALT), + ("BACKGROUND", (1, 1), (1, -1), PRINT_BG), + ("BACKGROUND", (1, 0), (1, 0), priority_color), + ("TEXTCOLOR", (0, 0), (0, -1), PRINT_MUTED), + ("TEXTCOLOR", (1, 1), (-1, -1), PRINT_TEXT), + ("TEXTCOLOR", (1, 0), (1, 0), PUNK_BG if priority_label in {"Low", "Medium"} else PUNK_WHITE), + ("FONTNAME", (0, 0), (0, -1), FONT_BOLD), + ("VALIGN", (0, 0), (-1, -1), "TOP"), + ("GRID", (0, 0), (-1, -1), 0.25, PRINT_BORDER), + ("LEFTPADDING", (0, 0), (-1, -1), 6), + ("RIGHTPADDING", (0, 0), (-1, -1), 6), + ("TOPPADDING", (0, 0), (-1, -1), 5), + ("BOTTOMPADDING", (0, 0), (-1, -1), 5), + ] + ) + ) + + return [ + Paragraph(self._escaped(f"ISSUE {index:02d}"), styles["Eyebrow"]), + self._toc_heading( + self._escaped(f"ISSUE {index:02d} - {finding.title} - {finding.file}"), + styles["Heading2"], + 1, + f"issue-{index}", + ), + Spacer(1, 0.1 * inch), + table, + Spacer(1, 0.18 * inch), + Paragraph("Issue", styles["Heading3"]), + Spacer(1, 0.08 * inch), + self._text_panel(finding.issue, styles), + Spacer(1, 0.20 * inch), + Paragraph("Recommendation", styles["Heading3"]), + Spacer(1, 0.08 * inch), + self._text_panel(finding.recommendation or "Not applicable.", styles), + Spacer(1, 0.18 * inch), + Paragraph("Validation steps", styles["Heading3"]), + Spacer(1, 0.08 * inch), + self._text_panel(self._validation_steps_text(finding.validation_steps), styles), + Spacer(1, 0.18 * inch), + Paragraph("Affected code", styles["Heading3"]), + self._context_table(finding, styles), + ] + + def _styles(self) -> dict[str, ParagraphStyle]: + styles = getSampleStyleSheet() + styles.add( + ParagraphStyle( + name="Project", + parent=styles["Heading2"], + alignment=TA_CENTER, + textColor=PUNK_WHITE, + spaceAfter=12, + ) + ) + styles.add( + ParagraphStyle( + name="MutedCenter", + parent=styles["BodyText"], + alignment=TA_CENTER, + textColor=PUNK_MUTED, + spaceAfter=6, + ) + ) + styles.add( + ParagraphStyle( + name="Eyebrow", + parent=styles["BodyText"], + fontSize=8, + leading=10, + textColor=PUNK_SECONDARY, + fontName=FONT_BOLD, + spaceAfter=4, + ) + ) + styles.add( + ParagraphStyle( + name="EyebrowCenter", + parent=styles["Eyebrow"], + alignment=TA_CENTER, + textColor=PUNK_GREEN, + spaceAfter=8, + ) + ) + styles.add( + ParagraphStyle( + name="CoverSubtitle", + parent=styles["BodyText"], + alignment=TA_CENTER, + textColor=PUNK_TEXT, + fontName=FONT_BOLD, + fontSize=20, + leading=26, + ) + ) + styles.add( + ParagraphStyle( + name="Body", + parent=styles["BodyText"], + leading=14, + textColor=PRINT_TEXT, + spaceAfter=8, + ) + ) + styles.add( + ParagraphStyle( + name="CodeLine", + parent=styles["BodyText"], + fontName=FONT_MONO, + fontSize=6.8, + leading=8, + textColor=PRINT_TEXT, + splitLongWords=1, + ) + ) + styles.add( + ParagraphStyle( + name="CodeNumber", + parent=styles["CodeLine"], + textColor=PRINT_MUTED, + alignment=TA_CENTER, + ) + ) + styles.add( + ParagraphStyle( + name="CodeLineHighlight", + parent=styles["CodeLine"], + textColor=PRINT_TEXT, + ) + ) + styles.add( + ParagraphStyle( + name="CodeNumberHighlight", + parent=styles["CodeNumber"], + textColor=PRINT_TEXT, + fontName=FONT_MONO, + ) + ) + styles.add( + ParagraphStyle( + name="PanelLabel", + parent=styles["Eyebrow"], + textColor=PUNK_SECONDARY, + ) + ) + styles.add( + ParagraphStyle( + name="PanelBody", + parent=styles["Body"], + backColor=PRINT_PANEL, + borderColor=PRINT_BORDER, + borderWidth=0.4, + borderPadding=8, + rightIndent=CONTENT_RIGHT_INDENT, + leading=14, + spaceAfter=8, + ) + ) + styles.add( + ParagraphStyle( + name="MetadataLabel", + parent=styles["BodyText"], + fontName=FONT_BOLD, + fontSize=8, + leading=10, + textColor=PUNK_GREEN, + ) + ) + styles.add( + ParagraphStyle( + name="MetadataValue", + parent=styles["BodyText"], + fontSize=9, + leading=11, + textColor=PUNK_TEXT, + spaceAfter=0, + ) + ) + styles.add( + ParagraphStyle( + name="CoverDisclaimerLabel", + parent=styles["MetadataLabel"], + fontSize=8, + leading=10, + textColor=PUNK_ORANGE, + ) + ) + styles.add( + ParagraphStyle( + name="CoverDisclaimer", + parent=styles["BodyText"], + fontSize=8.5, + leading=11, + textColor=PUNK_MUTED, + spaceAfter=0, + ) + ) + styles.add( + ParagraphStyle( + name="TocLevel0", + parent=styles["BodyText"], + fontName=FONT_BOLD, + fontSize=11, + leading=15, + leftIndent=0, + firstLineIndent=0, + spaceBefore=8, + textColor=PUNK_SECONDARY, + ) + ) + styles.add( + ParagraphStyle( + name="TocLevel1", + parent=styles["TocLevel0"], + fontName=FONT_REGULAR, + fontSize=9, + leading=12, + leftIndent=18, + firstLineIndent=0, + spaceBefore=4, + textColor=PRINT_MUTED, + ) + ) + styles.add( + ParagraphStyle( + name="IssueSummaryHeader", + parent=styles["BodyText"], + fontName=FONT_BOLD, + fontSize=9, + leading=11, + textColor=PUNK_WHITE, + ) + ) + styles.add( + ParagraphStyle( + name="IssueSummaryCell", + parent=styles["Body"], + fontSize=9, + leading=11, + textColor=PRINT_TEXT, + spaceAfter=0, + ) + ) + styles.add( + ParagraphStyle( + name="IssueSummaryId", + parent=styles["IssueSummaryCell"], + fontName=FONT_BOLD, + textColor=PUNK_SECONDARY, + ) + ) + styles.add( + ParagraphStyle( + name="IssueSummarySeverity", + parent=styles["IssueSummaryCell"], + alignment=TA_CENTER, + fontName=FONT_BOLD, + ) + ) + + styles["Title"].alignment = TA_CENTER + styles["Title"].fontSize = 28 + styles["Title"].leading = 32 + styles["Title"].textColor = PUNK_WHITE + styles["Title"].fontName = FONT_BOLD + styles["BodyText"].fontName = FONT_REGULAR + styles["Heading1"].fontName = FONT_BOLD + styles["Heading2"].fontName = FONT_BOLD + styles["Heading3"].fontName = FONT_BOLD + styles["Heading1"].textColor = PUNK_SECONDARY + styles["Heading2"].textColor = PUNK_SECONDARY + styles["Heading3"].textColor = PUNK_SECONDARY + styles["Heading1"].spaceBefore = 12 + styles["Heading1"].spaceAfter = 10 + styles["Heading2"].spaceBefore = 8 + styles["Heading2"].spaceAfter = 8 + styles["Heading3"].spaceBefore = 6 + styles["Heading3"].spaceAfter = 4 + + return styles + + @staticmethod + def _toc_heading(text: str, style: ParagraphStyle, level: int, key: str) -> Paragraph: + paragraph = Paragraph(text, style) + paragraph._saist_toc_level = level + paragraph._saist_toc_key = key + return paragraph + + def _document_title(self) -> str: + if self.project: + return f"SAIST report - {self.project}" + return "SAIST report" + + def _cover_title_markup(self) -> str: + if self.project: + project_line = self._escaped(self.project) + return f"AI generated code review of
{project_line}
by {self._rainbow_saist_markup()}" + return f"AI generated code review
by {self._rainbow_saist_markup()}" + + @staticmethod + def _rainbow_saist_markup() -> str: + return "".join(f'{letter}' for letter, color in zip("SAIST", RAINBOW_HEX)) + + def _model_text(self) -> str: + parts = [self.llm.model_vendor, self.llm.model_name] + return ": ".join(part for part in parts if part) + + @staticmethod + def _scm_adapter_from_args(args) -> str: + return getattr(args, "SCM", None) or "Not specified" + + @staticmethod + def _target_from_args(args) -> str: + scm = getattr(args, "SCM", None) + + if scm in {"filesystem", "git"}: + return getattr(args, "path", None) or "Not specified" + + if scm == "github": + repository = getattr(args, "repository", None) + pr = getattr(args, "pr", None) + if repository and pr: + return f"{repository} pull request #{pr}" + return repository or "Not specified" + + return getattr(args, "path", None) or getattr(args, "repository", None) or "Not specified" + + def _disclaimer_panel(self, styles: dict[str, ParagraphStyle]) -> Table: + table = Table( + [ + [Paragraph("Disclaimer", styles["CoverDisclaimerLabel"])], + [ + Paragraph( + """ + This report has not been created by the expert testing team at Punk Security. + + It is created by SAIST, an open-source tool developed by Punk Security that uses AI to analyze code and generate findings. The findings and recommendations in this report are generated by an AI model based on the input provided to SAIST. + """, + styles["CoverDisclaimer"], + ) + ], + ], + colWidths=[CONTENT_WIDTH], + hAlign="LEFT", + ) + table.setStyle( + TableStyle( + [ + ("BACKGROUND", (0, 0), (-1, -1), PUNK_BG), + ("BOX", (0, 0), (-1, -1), 0.45, PUNK_SECONDARY_LIGHT), + ("LINEABOVE", (0, 0), (-1, 0), 1.0, PUNK_ORANGE), + ("LEFTPADDING", (0, 0), (-1, -1), 12), + ("RIGHTPADDING", (0, 0), (-1, -1), 12), + ("TOPPADDING", (0, 0), (-1, 0), 9), + ("BOTTOMPADDING", (0, 0), (-1, 0), 2), + ("TOPPADDING", (0, 1), (-1, 1), 0), + ("BOTTOMPADDING", (0, 1), (-1, 1), 10), + ] + ) + ) + return table + + def _metadata_table(self, styles: dict[str, ParagraphStyle]) -> Table: + rows = [ + ["Date", getattr(self, "generated_at", "Not specified")], + ["Model", self._model_text() or "Not specified"], + ["SCM adapter", getattr(self, "scm_adapter", "Not specified")], + ["Target", getattr(self, "target", "Not specified")], + ] + table_rows = [ + [ + Paragraph(self._escaped(label), styles["MetadataLabel"]), + Paragraph(self._escaped(value), styles["MetadataValue"]), + ] + for label, value in rows + ] + + table = Table(table_rows, colWidths=[1.35 * inch, 5.35 * inch], hAlign="LEFT") + table_style = [ + ("BACKGROUND", (0, 0), (-1, -1), PUNK_BG), + ("BOX", (0, 0), (-1, -1), 0.45, PUNK_SECONDARY_LIGHT), + ("LINEABOVE", (0, 0), (-1, 0), 1.0, PUNK_BLUE), + ("VALIGN", (0, 0), (-1, -1), "MIDDLE"), + ("LEFTPADDING", (0, 0), (0, -1), 12), + ("RIGHTPADDING", (0, 0), (0, -1), 8), + ("LEFTPADDING", (1, 0), (1, -1), 4), + ("RIGHTPADDING", (1, 0), (1, -1), 12), + ("TOPPADDING", (0, 0), (-1, -1), 7), + ("BOTTOMPADDING", (0, 0), (-1, -1), 7), + ] + for row_index in range(len(table_rows) - 1): + table_style.append(("LINEBELOW", (0, row_index), (-1, row_index), 0.25, PUNK_SECONDARY_LIGHT)) + table.setStyle(TableStyle(table_style)) + return table + + def _issue_summary_table(self, styles: dict[str, ParagraphStyle]) -> Table: + rows = [ + [ + Paragraph("Issue ID", styles["IssueSummaryHeader"]), + Paragraph("Title", styles["IssueSummaryHeader"]), + Paragraph("File", styles["IssueSummaryHeader"]), + Paragraph("Severity", styles["IssueSummaryHeader"]), + ] + ] + + for index, finding in enumerate(self.findings, start=1): + priority_label, _ = self._priority(finding.priority) + rows.append( + [ + Paragraph(self._escaped(f"ISSUE {index:02d}"), styles["IssueSummaryId"]), + Paragraph(self._escaped(finding.title), styles["IssueSummaryCell"]), + Paragraph(self._escaped(finding.file), styles["IssueSummaryCell"]), + Paragraph(self._escaped(priority_label), styles["IssueSummarySeverity"]), + ] + ) + + if len(rows) == 1: + rows.append( + [ + Paragraph("-", styles["IssueSummaryCell"]), + Paragraph("No issues were provided.", styles["IssueSummaryCell"]), + Paragraph("-", styles["IssueSummaryCell"]), + Paragraph("-", styles["IssueSummaryCell"]), + ] + ) + + table = Table(rows, colWidths=[0.92 * inch, 2.45 * inch, 2.35 * inch, 0.98 * inch], hAlign="LEFT", repeatRows=1) + table_style = [ + ("BACKGROUND", (0, 0), (-1, 0), PUNK_SECONDARY), + ("TEXTCOLOR", (0, 0), (-1, 0), PUNK_WHITE), + ("FONTNAME", (0, 0), (-1, 0), FONT_BOLD), + ("BACKGROUND", (0, 1), (-1, -1), PRINT_BG), + ("GRID", (0, 0), (-1, -1), 0.25, PRINT_BORDER), + ("VALIGN", (0, 0), (-1, -1), "MIDDLE"), + ("LEFTPADDING", (0, 0), (-1, -1), 8), + ("RIGHTPADDING", (0, 0), (-1, -1), 8), + ("TOPPADDING", (0, 0), (-1, -1), 7), + ("BOTTOMPADDING", (0, 0), (-1, -1), 7), + ("ALIGN", (3, 1), (3, -1), "CENTER"), + ] + + for row_index, finding in enumerate(self.findings, start=1): + if row_index % 2 == 0: + table_style.append(("BACKGROUND", (0, row_index), (-1, row_index), PRINT_PANEL)) + priority_label, priority_color = self._priority(finding.priority) + table_style.extend( + [ + ("BACKGROUND", (3, row_index), (3, row_index), priority_color), + ( + "TEXTCOLOR", + (3, row_index), + (3, row_index), + PUNK_BG if priority_label in {"Low", "Medium"} else PUNK_WHITE, + ), + ("FONTNAME", (3, row_index), (3, row_index), FONT_BOLD), + ] + ) + + table.setStyle(TableStyle(table_style)) + return table + + def _text_panel(self, text: str, styles: dict[str, ParagraphStyle]) -> Paragraph: + return Paragraph(self._paragraph_markup(text), styles["PanelBody"]) + + @staticmethod + def _validation_steps_text(validation_steps: list[str]) -> str: + if not validation_steps: + return "Not provided." + return "\n".join(f"{index}. {step}" for index, step in enumerate(validation_steps, start=1)) + + def _context_table(self, finding: FindingContext, styles: dict[str, ParagraphStyle]) -> Table: + lines = finding.context.splitlines() if finding.context else [""] + rows = [] + + for offset, line in enumerate(lines): + line_number = finding.context_start + offset + highlighted = line_number == finding.line_number + rows.append( + [ + str(line_number), + Paragraph(self._code_markup(line), styles["CodeLineHighlight" if highlighted else "CodeLine"]), + ] + ) + + table = Table(rows, colWidths=[CODE_LINE_NUMBER_WIDTH, CONTENT_WIDTH - CODE_LINE_NUMBER_WIDTH], hAlign="LEFT", repeatRows=0) + table_style = [ + ("BACKGROUND", (0, 0), (-1, -1), PRINT_CODE_BG), + ("BACKGROUND", (0, 0), (0, -1), PRINT_PANEL_ALT), + ("GRID", (0, 0), (-1, -1), 0.2, PRINT_BORDER), + ("VALIGN", (0, 0), (-1, -1), "TOP"), + ("ALIGN", (0, 0), (0, -1), "RIGHT"), + ("FONTNAME", (0, 0), (0, -1), FONT_MONO), + ("FONTSIZE", (0, 0), (0, -1), 6.8), + ("LEADING", (0, 0), (0, -1), 8), + ("TEXTCOLOR", (0, 0), (0, -1), PRINT_MUTED), + ("LEFTPADDING", (0, 0), (-1, -1), 4), + ("RIGHTPADDING", (0, 0), (-1, -1), 4), + ("TOPPADDING", (0, 0), (-1, -1), 2), + ("BOTTOMPADDING", (0, 0), (-1, -1), 2), + ] + + for index, _ in enumerate(lines): + line_number = finding.context_start + index + if line_number == finding.line_number: + table_style.extend( + [ + ("BACKGROUND", (0, index), (-1, index), PRINT_CODE_HIGHLIGHT), + ("TEXTCOLOR", (0, index), (-1, index), PRINT_TEXT), + ("FONTNAME", (0, index), (0, index), "Courier-Bold"), + ] + ) + + table.setStyle(TableStyle(table_style)) + return table + + @staticmethod + def _logo_path() -> str | None: + path = os.path.abspath( + os.path.join( + os.path.dirname(__file__), + "assets", + "PunkSecurityLogo.png", + ) + ) + return path if os.path.exists(path) else None + + @staticmethod + def _draw_page(canvas, doc): + width, height = A4 + canvas.saveState() + page_is_cover = doc.page == 1 + canvas.setFillColor(PUNK_SECONDARY if page_is_cover else PRINT_BG) + canvas.rect(0, 0, width, height, stroke=0, fill=1) + canvas.setFillColor(PUNK_SECONDARY) + canvas.rect(0, height - 0.38 * inch, width, 0.38 * inch, stroke=0, fill=1) + ReportLabPdf._draw_rainbow_bar(canvas, width, height - 0.40 * inch) + + if not page_is_cover: + canvas.setFillColor(PUNK_WHITE) + canvas.setFont(FONT_REGULAR, 8) + canvas.drawString(0.65 * inch, height - 0.25 * inch, "SAIST AI generated code review") + canvas.drawRightString(width - 0.65 * inch, 0.35 * inch, f"Page {doc.page}") + ReportLabPdf._draw_rainbow_bar(canvas, width, 0.22 * inch) + + canvas.restoreState() + + @staticmethod + def _draw_rainbow_bar(canvas, width: float, y: float): + segment_width = width / len(RAINBOW) + for index, color in enumerate(RAINBOW): + canvas.setFillColor(color) + canvas.rect(index * segment_width, y, segment_width + 1, 0.035 * inch, stroke=0, fill=1) + + @staticmethod + def _escaped(value) -> str: + return escape(str(value or "")) + + def _paragraph_markup(self, value: str) -> str: + escaped = self._escaped(value) + return escaped.replace("\n\n", "

").replace("\n", "
") + + def _code_markup(self, value: str) -> str: + line = str(value or " ") + wrapped = textwrap.wrap( + line, + width=CODE_WRAP_CHARS, + break_long_words=True, + break_on_hyphens=False, + replace_whitespace=False, + drop_whitespace=False, + ) or [" "] + return "
".join(self._escaped(part).replace(" ", " ") for part in wrapped) + + @staticmethod + def _priority(priority: int): + if priority > 8: + return "Critical", PUNK_RED + if priority > 7: + return "High", PUNK_RED + if priority > 4: + return "Medium", PUNK_ORANGE + return "Low", PUNK_BLUE diff --git a/saist/latex/tex/image/PunkSecurityLogo.png b/saist/reportlab_pdf/assets/PunkSecurityLogo.png similarity index 100% rename from saist/latex/tex/image/PunkSecurityLogo.png rename to saist/reportlab_pdf/assets/PunkSecurityLogo.png diff --git a/saist/reportlab_pdf/fonts/OFL.txt b/saist/reportlab_pdf/fonts/OFL.txt new file mode 100644 index 0000000..8fd08fa --- /dev/null +++ b/saist/reportlab_pdf/fonts/OFL.txt @@ -0,0 +1,93 @@ +Copyright 2019 The Pathway Extreme Project Authors (https://github.com/etunni/Pathway-Variable-Font) + +This Font Software is licensed under the SIL Open Font License, Version 1.1. +This license is copied below, and is also available with a FAQ at: +https://scripts.sil.org/OFL + + +----------------------------------------------------------- +SIL OPEN FONT LICENSE Version 1.1 - 26 February 2007 +----------------------------------------------------------- + +PREAMBLE +The goals of the Open Font License (OFL) are to stimulate worldwide +development of collaborative font projects, to support the font creation +efforts of academic and linguistic communities, and to provide a free and +open framework in which fonts may be shared and improved in partnership +with others. + +The OFL allows the licensed fonts to be used, studied, modified and +redistributed freely as long as they are not sold by themselves. The +fonts, including any derivative works, can be bundled, embedded, +redistributed and/or sold with any software provided that any reserved +names are not used by derivative works. The fonts and derivatives, +however, cannot be released under any other type of license. The +requirement for fonts to remain under this license does not apply +to any document created using the fonts or their derivatives. + +DEFINITIONS +"Font Software" refers to the set of files released by the Copyright +Holder(s) under this license and clearly marked as such. This may +include source files, build scripts and documentation. + +"Reserved Font Name" refers to any names specified as such after the +copyright statement(s). + +"Original Version" refers to the collection of Font Software components as +distributed by the Copyright Holder(s). + +"Modified Version" refers to any derivative made by adding to, deleting, +or substituting -- in part or in whole -- any of the components of the +Original Version, by changing formats or by porting the Font Software to a +new environment. + +"Author" refers to any designer, engineer, programmer, technical +writer or other person who contributed to the Font Software. + +PERMISSION & CONDITIONS +Permission is hereby granted, free of charge, to any person obtaining +a copy of the Font Software, to use, study, copy, merge, embed, modify, +redistribute, and sell modified and unmodified copies of the Font +Software, subject to the following conditions: + +1) Neither the Font Software nor any of its individual components, +in Original or Modified Versions, may be sold by itself. + +2) Original or Modified Versions of the Font Software may be bundled, +redistributed and/or sold with any software, provided that each copy +contains the above copyright notice and this license. These can be +included either as stand-alone text files, human-readable headers or +in the appropriate machine-readable metadata fields within text or +binary files as long as those fields can be easily viewed by the user. + +3) No Modified Version of the Font Software may use the Reserved Font +Name(s) unless explicit written permission is granted by the corresponding +Copyright Holder. This restriction only applies to the primary font name as +presented to the users. + +4) The name(s) of the Copyright Holder(s) or the Author(s) of the Font +Software shall not be used to promote, endorse or advertise any +Modified Version, except to acknowledge the contribution(s) of the +Copyright Holder(s) and the Author(s) or with their explicit written +permission. + +5) The Font Software, modified or unmodified, in part or in whole, +must be distributed entirely under this license, and must not be +distributed under any other license. The requirement for fonts to +remain under this license does not apply to any document created +using the Font Software. + +TERMINATION +This license becomes null and void if any of the above conditions are +not met. + +DISCLAIMER +THE FONT SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, +EXPRESS OR IMPLIED, INCLUDING BUT NOT LIMITED TO ANY WARRANTIES OF +MERCHANTABILITY, FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT +OF COPYRIGHT, PATENT, TRADEMARK, OR OTHER RIGHT. IN NO EVENT SHALL THE +COPYRIGHT HOLDER BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER LIABILITY, +INCLUDING ANY GENERAL, SPECIAL, INDIRECT, INCIDENTAL, OR CONSEQUENTIAL +DAMAGES, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING +FROM, OUT OF THE USE OR INABILITY TO USE THE FONT SOFTWARE OR FROM +OTHER DEALINGS IN THE FONT SOFTWARE. diff --git a/saist/reportlab_pdf/fonts/PathwayExtreme.ttf b/saist/reportlab_pdf/fonts/PathwayExtreme.ttf new file mode 100644 index 0000000..d18c67e Binary files /dev/null and b/saist/reportlab_pdf/fonts/PathwayExtreme.ttf differ diff --git a/saist/scm/__init__.py b/saist/scm/__init__.py index 2d6378f..41536b8 100644 --- a/saist/scm/__init__.py +++ b/saist/scm/__init__.py @@ -1,5 +1,4 @@ -from os import PathLike -from typing import TypedDict +from typing import Callable, TypedDict from .adapters import BaseScmAdapter @@ -48,5 +47,39 @@ async def read_file_contents(self, filename: str): """ return await self.adapter.get_file_contents(filename) + async def list_files(self) -> list[str]: + """ + Lists all files available to the scanner, relative to the source root. + """ + return await self.adapter.list_files() + + async def regex_search( + self, + pattern: str, + file_pattern: str = "**/*", + max_results: int = 100, + ) -> list[dict[str, str | int]]: + """ + Searches files available to the scanner using a Python regular expression. + + Args: + pattern: Python regular expression to search for. Inline flags like (?i) are supported. + file_pattern: Optional glob for limiting files, for example **/*.py. + max_results: Maximum number of matches to return. + """ + return await self.adapter.regex_search(pattern, file_pattern, max_results) + + def tool_functions(self) -> list[Callable]: + """ + Returns the SCM helper functions exposed to the LLM. + """ + return [self.read_file_contents, self.list_files, self.regex_search] + + def detect_prompt(self) -> str: + return self.adapter.detect_prompt() + + def summary_prompt(self) -> str: + return self.adapter.summary_prompt() + def create_review(self, comment, review_comments, request_changes): self.adapter.create_review(comment, review_comments, request_changes) diff --git a/saist/scm/adapters/__init__.py b/saist/scm/adapters/__init__.py index 81fd3e7..c22441d 100644 --- a/saist/scm/adapters/__init__.py +++ b/saist/scm/adapters/__init__.py @@ -1,8 +1,19 @@ from abc import ABCMeta, abstractmethod -from os import PathLike +import fnmatch +from pathlib import PurePosixPath +import re class BaseScmAdapter(metaclass=ABCMeta): + DETECT_PROMPT = "" + SUMMARY_PROMPT = "" + + def detect_prompt(self) -> str: + return self.DETECT_PROMPT + + def summary_prompt(self) -> str: + return self.SUMMARY_PROMPT + @abstractmethod def create_review(self, comment, review_comments, request_changes): """ @@ -26,3 +37,76 @@ def get_changed_files(self): @abstractmethod async def get_file_contents(self, file_path: str): pass + + async def list_files(self) -> list[str]: + """ + Lists every file path available to this adapter, relative to the project root. + Adapters that cannot enumerate files can return an empty list. + """ + return [] + + async def regex_search( + self, + pattern: str, + file_pattern: str = "**/*", + max_results: int = 100, + ) -> list[dict[str, str | int]]: + """ + Searches available UTF-8 files using a Python regular expression. + + Args: + pattern: Python regular expression to search for. Inline flags like (?i) are supported. + file_pattern: Optional glob for limiting files, for example **/*.py. + max_results: Maximum number of matches to return. + + Returns: + A list of matches with filename, line_number, column, match, and line fields. + """ + if max_results <= 0: + return [] + + try: + regex = re.compile(pattern) + except re.error as e: + return [{"error": f"Invalid regex: {e}", "pattern": pattern}] + + results = [] + for filename in await self.list_files(): + if not self._file_matches_pattern(filename, file_pattern): + continue + + try: + contents = await self.get_file_contents(filename) + except Exception: + continue + + if contents is None: + continue + + for line_number, line in enumerate(contents.splitlines(), start=1): + for match in regex.finditer(line): + results.append( + { + "filename": filename, + "line_number": line_number, + "column": match.start() + 1, + "match": match.group(0), + "line": line, + } + ) + if len(results) >= max_results: + return results + + return results + + @staticmethod + def _file_matches_pattern(filename: str, file_pattern: str) -> bool: + if not file_pattern or file_pattern == "**/*": + return True + + path = PurePosixPath(filename) + patterns = [file_pattern] + if file_pattern.startswith("**/"): + patterns.append(file_pattern[3:]) + + return any(path.match(pattern) or fnmatch.fnmatch(filename, pattern) for pattern in patterns) diff --git a/saist/scm/adapters/filesystem.py b/saist/scm/adapters/filesystem.py index dacdb9d..86f6a6f 100644 --- a/saist/scm/adapters/filesystem.py +++ b/saist/scm/adapters/filesystem.py @@ -16,19 +16,39 @@ logger = logging.getLogger(__name__) class FilesystemAdapter(BaseScmAdapter): + DETECT_PROMPT = """ +You are performing a penetration test style review across the entire application codebase. +The supplied input is one file from the application. It may be represented as a unified diff from an empty file, but you should treat it as application code in scope for a whole-codebase security assessment. +Use tools aggressively to map routes, controllers, models, middleware, policies, serializers, templates, jobs, and configuration before deciding what is exploitable. +Trace attacker-controlled input from entrypoint to sink and trace authorization decisions from identity source to protected action. +Look for business logic vulnerabilities across files: cross-tenant data access, horizontal/vertical privilege escalation, order/payment/state-machine manipulation, invitation or password-reset abuse, webhook forgery, unsafe admin actions, and background jobs that trust user-controlled state. +Report vulnerabilities that are present in the application, even when the exploit depends on interactions across multiple files. +Avoid one-file lint findings unless that file alone proves a reachable vulnerability. +""" + + SUMMARY_PROMPT = """ +This summary is for a penetration test style review across the entire application codebase. +Summarize exploitable application risks, affected trust boundaries, likely business impact, and the highest-impact fixes. Do not summarize generic best practices. +""" + async def get_file_contents(self, filename: str): logger.debug(f"get_file_contents: reading file {filename} under {self.compare_path}") - filename = os.path.join(self.compare_path, filename) try: - if not pathlib.Path(filename).is_relative_to(self.compare_path): - raise Exception(f"Tried to access file outside of the root: {filename}") + file_path = self._resolve_under_root(filename) - async with aiofiles.open(filename, mode='r', encoding='utf-8') as f: - contents = await f.read() + async with aiofiles.open(file_path, mode="rb") as f: + data = await f.read() - return contents + return data.decode("utf-8") + except UnicodeDecodeError: + logger.debug(f"get_file_contents: file is not valid UTF-8, skipping: {filename}") + return None + except (FileNotFoundError, IsADirectoryError): + logger.debug(f"get_file_contents: file does not exist or is not readable: {filename}") + return None except Exception as e: - logging.warn(f"ERR: {e}") + logger.warning(f"get_file_contents: could not read {filename}: {e}") + return None def __init__(self, compare_path: PathLike[AnyStr] | str, base_path: Optional[PathLike[AnyStr] | str] = None): self.base_path = base_path @@ -40,6 +60,28 @@ def create_review(self, comment, review_comments, request_changes): def get_changed_files(self) -> list[File]: return list(self._iter_changed_files()) + async def list_files(self) -> list[str]: + """ + Lists all regular files under the comparison path. + Paths are relative to the comparison path and use POSIX separators. + """ + root = self._root_path() + files = [] + + for path in root.rglob("*"): + try: + resolved_path = path.resolve() + except OSError as e: + logger.warning(f"Could not resolve file path {path}: {e}") + continue + + if not resolved_path.is_relative_to(root) or not resolved_path.is_file(): + continue + + files.append(path.relative_to(root).as_posix()) + + return sorted(files) + def _iter_changed_files(self) -> list[File]: logger.debug(f"Iterate changed files: base:{self.base_path}, compare:{self.compare_path}") @@ -80,3 +122,15 @@ def _iter_changed_files(self) -> list[File]: def likely(): # TODO: implement proper logic return True + + def _root_path(self) -> pathlib.Path: + return pathlib.Path(self.compare_path).resolve() + + def _resolve_under_root(self, filename: str) -> pathlib.Path: + root = self._root_path() + file_path = (root / filename).resolve() + + if not file_path.is_relative_to(root): + raise Exception(f"Tried to access file outside of the root: {filename}") + + return file_path diff --git a/saist/scm/adapters/git.py b/saist/scm/adapters/git.py index d0a7ae7..bbf5b61 100644 --- a/saist/scm/adapters/git.py +++ b/saist/scm/adapters/git.py @@ -1,6 +1,7 @@ import io import logging from os import PathLike +from pathlib import PurePosixPath from typing import Optional, Iterator from scm import BaseScmAdapter, File @@ -11,11 +12,34 @@ logger = logging.getLogger(__name__) class GitAdapter(BaseScmAdapter): + DETECT_PROMPT = """ +You are analyzing a diff of code that needs security review. +The supplied input is a single file's unified diff from a git comparison. +Focus on exploitable vulnerabilities introduced, exposed, or materially changed by this diff. +Use tools to retrieve the full file and related files to validate whether the changed code is reachable, attacker-controlled, and crosses a security boundary. +Report only vulnerabilities anchored to changed lines in the original diff. Do not report pre-existing best-practice issues unless the diff makes them exploitable or materially worse. +For business logic changes, inspect surrounding authorization, state transition, tenancy, payment, invitation, webhook, or admin-flow code before deciding. +""" + + SUMMARY_PROMPT = """ +This summary is for a diff-based code security review. +Summarize exploitable risks introduced or changed by the supplied diff, including business impact and affected security boundaries. Do not summarize generic best practices. +""" + async def get_file_contents(self, file_path: str): logger.debug(f"file_get_contents: Reading {file_path}") - targetfile = self.compare_commit.tree / file_path - with io.BytesIO(targetfile.data_stream.read()) as f: - return f.read().decode('utf-8') + clean_path = self._clean_repo_path(file_path) + if clean_path is None: + logger.warning(f"get_file_contents: rejected path outside repository: {file_path}") + return None + + try: + targetfile = self.compare_commit.tree / clean_path + with io.BytesIO(targetfile.data_stream.read()) as f: + return f.read().decode('utf-8') + except (KeyError, ValueError, UnicodeDecodeError) as e: + logger.warning(f"get_file_contents: could not read {clean_path}: {e}") + return None def __init__(self, repo_path: Optional[PathLike]=None, base_branch: Optional[str] = None, compare_branch: Optional[str] = None, base_commit: Optional[str] = None, compare_commit: Optional[str] = None): self.repo_path = repo_path @@ -60,18 +84,30 @@ def _repo(self) -> Repo: return repo def _iter_diffs(self) -> Iterator[File]: - diffs = self.compare_commit.diff(self.base_commit, create_patch=True) + diffs = self.base_commit.diff(self.compare_commit, create_patch=True) for diff in diffs: + filename = diff.b_path or diff.a_path match diff.diff: case bytes(data): - yield File(filename=diff.a_path, patch=data.decode('UTF-8')) + yield File(filename=filename, patch=data.decode('UTF-8')) case data: - yield File(filename=diff.a_path, patch=data) + yield File(filename=filename, patch=data) def get_changed_files(self) -> list[File]: return list([f for f in self._iter_diffs() if f['filename'] != None]) + async def list_files(self) -> list[str]: + """ + Lists all blob paths at the comparison commit. + """ + files = [] + for item in self.compare_commit.tree.traverse(): + if item.type == "blob": + files.append(item.path) + + return sorted(files) + def create_review(self, comment, review_comments, request_changes): write_findings(comment,review_comments,request_changes) @@ -79,3 +115,12 @@ def create_review(self, comment, review_comments, request_changes): def likely(): # TODO: determine logic return True + + @staticmethod + def _clean_repo_path(file_path: str) -> Optional[str]: + path = PurePosixPath(str(file_path).replace("\\", "/")) + + if path.is_absolute() or ".." in path.parts: + return None + + return path.as_posix() diff --git a/saist/scm/adapters/github.py b/saist/scm/adapters/github.py index 783c14d..3d2d021 100644 --- a/saist/scm/adapters/github.py +++ b/saist/scm/adapters/github.py @@ -8,6 +8,7 @@ from requests import Request from requests.auth import AuthBase +from requests.exceptions import HTTPError from . import BaseScmAdapter @@ -24,6 +25,19 @@ def __call__(self, r: Request): class Github(BaseScmAdapter): API_BASE_URL = os.getenv("GITHUB_API_URL", "https://api.github.com") + DETECT_PROMPT = """ +You are analyzing a diff of code that needs security review. +The supplied input is a single file's unified diff from a GitHub pull request. +Focus on exploitable vulnerabilities introduced, exposed, or materially changed by this pull request. +Use tools to retrieve the full file and related files to validate whether the changed code is reachable, attacker-controlled, and crosses a security boundary. +Report only vulnerabilities anchored to changed lines in the original diff. Do not report pre-existing best-practice issues unless the pull request makes them exploitable or materially worse. +For business logic changes, inspect surrounding authorization, state transition, tenancy, payment, invitation, webhook, or admin-flow code before deciding. +""" + + SUMMARY_PROMPT = """ +This summary is for a GitHub pull request diff-based code security review. +Summarize exploitable risks introduced or changed by the pull request, including business impact and affected security boundaries. Do not summarize generic best practices. +""" def __init__(self, github_token, repo, pr_number): self.github_token = github_token @@ -61,7 +75,13 @@ async def get_file_contents(self, file_path: PathLike[str] | str): clean_url = f"{self.API_BASE_URL}/repos/{self.repo}/contents/{clean_path}?ref={self.commit_sha}" res = self.requests.get(clean_url, headers={"Content-Type": "application/vnd.github.object+json"}) - res.raise_for_status() + try: + res.raise_for_status() + except HTTPError: + if res.status_code == 404: + logger.warning(f"get_file_contents {clean_path}: file does not exist at PR head; skipping.") + return None + raise json_res = res.json() @@ -130,7 +150,7 @@ def get_changed_files(self): files_page = resp.json() if not files_page: break - all_files.extend(files_page) + all_files.extend(file for file in files_page if file.get("status") != "removed") page += 1 return all_files diff --git a/saist/shell/__init__.py b/saist/shell/__init__.py index ea82968..a072718 100644 --- a/saist/shell/__init__.py +++ b/saist/shell/__init__.py @@ -35,7 +35,17 @@ def __init__(self, llm: BaseLlmAdapter, scm: Scm, findings: list[Finding]): self.findings = findings self.original_findings = findings self.should_stop = False - self.agent = llm.generate_agent(self.PROMPT, [self.stop, self.get_findings, self.update_findings, self.reset_findings, self.reset_chat, scm.read_file_contents]) + self.agent = llm.generate_agent( + self.PROMPT, + [ + self.stop, + self.get_findings, + self.update_findings, + self.reset_findings, + self.reset_chat, + *scm.tool_functions(), + ], + ) self.new_messages = None self.console = Console() diff --git a/saist/util/argparsing.py b/saist/util/argparsing.py index 1cdcace..8933437 100644 --- a/saist/util/argparsing.py +++ b/saist/util/argparsing.py @@ -1,8 +1,9 @@ -import argparse -from os import linesep, environ, cpu_count -import sys -from shutil import which -from dotenv import load_dotenv +import argparse +from os import linesep, environ, cpu_count +import sys +from dotenv import load_dotenv +from llm.adapters import THINKING_CHOICES +from util.skills import DEFAULT_SKILL_MAX_BYTES, DEFAULT_SKILL_SAMPLE_BYTES, DEFAULT_SKILL_SAMPLE_FILES, DEFAULT_SKILLS_PATH load_dotenv(".env") @@ -45,11 +46,13 @@ def error(self, message): {runtime} --llm openai --llm-model gpt4o git {linesep} > Scan local code folder with ollama and get interactive shell {runtime} --llm ollama --interactive filesystem {linesep} -> Scan local code folder with anthropic and get web server findings -{runtime} --llm anthropic --interactive filesystem {linesep} -{linesep} -""", -) +> Scan local code folder with anthropic and get web server findings +{runtime} --llm anthropic --interactive filesystem {linesep} +> Generate project-specific analysis skills for future scans +{runtime} --llm openai --generate-skills filesystem {linesep} +{linesep} +""", +) class EnvDefault(argparse.Action): def __init__(self, envvar, required=True, default=None, **kwargs): @@ -93,13 +96,13 @@ def __call__(self, parser, namespace, values, option_string=None): git_parser.add_argument("--ref-to-compare", type=str, help = "Git ref to compare to", envvar="SAIST_GIT_COMPARE_REF", action=EnvDefault, default="HEAD" ) -git_parser.add_argument( - "--commit-for-compare", type=str, help = "Git commit to compare from (preferred over REF if set)", - envvar="SAIST_GIT_BASE_COMMIT", action=EnvDefault - ) -git_parser.add_argument("--commit-to-compare", type=str, help = "Git commit to compare to (preferred over REF if set)", - envvar="SAIST_GIT_COMPARE_COMMIT", action=EnvDefault - ) +git_parser.add_argument( + "--commit-for-compare", type=str, help = "Git commit to compare from (preferred over REF if set)", + envvar="SAIST_GIT_BASE_COMMIT", action=EnvDefault, required=False + ) +git_parser.add_argument("--commit-to-compare", type=str, help = "Git commit to compare to (preferred over REF if set)", + envvar="SAIST_GIT_COMPARE_COMMIT", action=EnvDefault, required=False + ) ### GITHUB github_parser = SCM_subparsers.add_parser("github", help = "Scan a github PR") @@ -119,7 +122,7 @@ def __call__(self, parser, namespace, values, option_string=None): parser.add_argument( "--llm", type=str, - choices=["anthropic", "bedrock", "deepseek", "gemini", "ollama", "openai", "faike"], + choices=["anthropic", "azure-foundry", "bedrock", "deepseek", "gemini", "ollama", "openai", "faike"], required=True, action=EnvDefault, envvar="SAIST_LLM" @@ -135,30 +138,73 @@ def __call__(self, parser, namespace, values, option_string=None): envvar="SAIST_LLM_MODEL", action=EnvDefault, required=False ) -parser.add_argument( - "--llm-rate-limit", help = "Max requests per second", envvar="SAIST_LLM_RATE_LIMIT", - action=EnvDefault, required=False, type=int, default = 10 - ) - -parser.add_argument( - "--ollama-base-uri", type=str, help = "Base uri of ollama", - envvar="SAIST_OLLAMA_BASE_URI", action=EnvDefault, default = "http://localhost:11434" - ) - -parser.add_argument( - "--openai-base-uri", type=str, help = "Base uri of openai to use any compatable service", - envvar="SAIST_OPENAI_BASE_URI", action=EnvDefault, required=False - ) - -parser.add_argument( +parser.add_argument( + "--llm-rate-limit", help = "Max requests per second", envvar="SAIST_LLM_RATE_LIMIT", + action=EnvDefault, required=False, type=int, default = 10 + ) + +parser.add_argument( + "--iterations", + help="Number of tool-driven filesystem analysis passes to run when --deep is not set", + envvar="SAIST_ITERATIONS", + action=EnvDefault, + required=False, + type=int, + default=1, +) + +parser.add_argument( + "--thinking", + help="LLM thinking effort for providers supported by pydantic-ai", + choices=THINKING_CHOICES, + envvar="SAIST_THINKING", + action=EnvDefault, + required=False, + default="medium", +) + +parser.add_argument( + "--ollama-base-uri", type=str, help = "Base uri of ollama", + envvar="SAIST_OLLAMA_BASE_URI", action=EnvDefault, default = "http://localhost:11434" + ) + +parser.add_argument( + "--openai-base-uri", "--open-ai-baseuri", type=str, help = "Base uri of openai to use any compatable service", + envvar="SAIST_OPENAI_BASE_URI", action=EnvDefault, required=False + ) + +parser.add_argument( + "--azure-openai-endpoint", + type=str, + help="Azure AI Foundry or Azure OpenAI endpoint (can be set with AZURE_OPENAI_ENDPOINT)", + envvar="AZURE_OPENAI_ENDPOINT", + action=EnvDefault, + required=False, +) + +parser.add_argument( + "--azure-openai-api-version", + type=str, + help="Azure OpenAI API version for non-v1 endpoints (can be set with OPENAI_API_VERSION)", + envvar="OPENAI_API_VERSION", + action=EnvDefault, + required=False, +) + +parser.add_argument( "--interactive", help = "Spawn an interactive prompt with the LLM at the end", required=False, action='store_true' ) -parser.add_argument( - "--disable-tools", help="Disable usage of tools during code analysis (this is a good cost saving)", - required=False, action="store_true" - ) +parser.add_argument( + "--disable-tools", help="Disable usage of tools during code analysis (this is a good cost saving)", + required=False, action="store_true" + ) + +parser.add_argument( + "--deep", help="For filesystem scans, analyze each file individually instead of using the tool-driven whole-application scan", + required=False, action="store_true" + ) parser.add_argument( "--web", help = "Launch a web server to display findings", @@ -190,39 +236,64 @@ def __call__(self, parser, namespace, values, option_string=None): envvar="SAIST_CSV_PATH", action=EnvDefault, required=False, default="results.csv" ) -parser.add_argument( - "--tex", help = "Write results of TeX file", - required=False, action='store_true' - ) - -parser.add_argument( - "--tex-filename", type=str, help = "Filename of TeX file", - envvar="SAIST_TEX_FILENAME", action=EnvDefault, required=False, default="report.tex" - ) - -parser.add_argument( - "--pdf", help = "Write results of PDF report", - required=False, action='store_true' - ) - -parser.add_argument( - "--pdf-filename", type=str, help = "Filename of PDF report", - envvar="SAIST_PDF_FILENAME", action=EnvDefault, required=False, default="report.pdf" - ) +parser.add_argument( + "--pdf", help = "Write results of PDF report", + required=False, action='store_true' + ) + +parser.add_argument( + "--pdf-filename", type=str, help = "Filename of PDF report", + envvar="SAIST_PDF_FILENAME", action=EnvDefault, required=False, default="report.pdf" + ) parser.add_argument( "--disable-caching", help = "Disable local caching of results", action='store_true', required=False ) -parser.add_argument( - "--cache-folder", type=str, help = "Folder name for local caching", - envvar="SAIST_CACHE_FOLDER", action=EnvDefault, required=False, default="SAISTCache" - ) - -parser.add_argument( - "--project-name", type=str, help = "Project name for pdf output", - envvar="SAIST_PROJECT_NAME", action=EnvDefault, required=False, default=None +parser.add_argument( + "--cache-folder", type=str, help = "Folder name for local caching", + envvar="SAIST_CACHE_FOLDER", action=EnvDefault, required=False, default="SAISTCache" + ) + +parser.add_argument( + "--skills-path", type=str, help = "Folder containing SAIST analysis skill Markdown files", + envvar="SAIST_SKILLS_PATH", action=EnvDefault, required=False, default=DEFAULT_SKILLS_PATH + ) + +parser.add_argument( + "--disable-skills", help = "Do not load SAIST analysis skill files during scanning", + required=False, action='store_true' + ) + +parser.add_argument( + "--generate-skills", help = "Generate SAIST analysis skill files for this project and exit", + required=False, action='store_true' + ) + +parser.add_argument( + "--overwrite-skills", help = "Replace existing skill files when used with --generate-skills", + required=False, action='store_true' + ) + +parser.add_argument( + "--skills-max-bytes", type=int, help = "Maximum total bytes of skill guidance to load into analysis prompts", + envvar="SAIST_SKILLS_MAX_BYTES", action=EnvDefault, required=False, default=DEFAULT_SKILL_MAX_BYTES + ) + +parser.add_argument( + "--skills-sample-files", type=int, help = "Maximum number of project files to sample when generating skills", + envvar="SAIST_SKILLS_SAMPLE_FILES", action=EnvDefault, required=False, default=DEFAULT_SKILL_SAMPLE_FILES + ) + +parser.add_argument( + "--skills-sample-bytes", type=int, help = "Maximum bytes to read from each sampled file when generating skills", + envvar="SAIST_SKILLS_SAMPLE_BYTES", action=EnvDefault, required=False, default=DEFAULT_SKILL_SAMPLE_BYTES + ) + +parser.add_argument( + "--project-name", type=str, help = "Project name for pdf output", + envvar="SAIST_PROJECT_NAME", action=EnvDefault, required=False, default=None ) parser.add_argument( @@ -256,8 +327,8 @@ def __call__(self, parser, namespace, values, option_string=None): help="-v for verbose, -vv for extra verbose", ) -def parse_args(): - args = parser.parse_args() +def parse_args(): + args = parser.parse_args() if args.llm == "bedrock" and args.llm_api_key: parser.error(f"Do not provide an API key for bedrock, use AWS ENV variables https://docs.aws.amazon.com/cli/v1/userguide/cli-configure-envvars.html") @@ -265,16 +336,30 @@ def parse_args(): if args.llm == "bedrock" and args.interactive: parser.error("Sorry, we dont support interactive mode with bedrock as AWS tool calling is a bit broken") - if args.llm not in [ "ollama", "bedrock", "faike" ] and args.llm_api_key is None: - parser.error(f"You must provide an api key with --llm-api-key if using {args.llm}") + if args.llm not in [ "azure-foundry", "ollama", "bedrock", "faike" ] and args.llm_api_key is None: + parser.error(f"You must provide an api key with --llm-api-key if using {args.llm}") if args.llm == "ollama" and args.interactive: parser.error(f"You cannot use the interactive shell with ollama currently") - if args.llm == "faike" and args.interactive: - parser.error("Faike LLM: Certified non-existent AI doesn't support interactive mode") - - if args.pdf and which("latexmk") == None: - parser.error("Unable to find 'latexmk' binary in $PATH needed for PDF report building, cannot use --pdf flag") - - return args + if args.llm == "faike" and args.interactive: + parser.error("Faike LLM: Certified non-existent AI doesn't support interactive mode") + + if args.generate_skills and args.SCM == "poem": + parser.error("Cannot generate SAIST skills while using the poem command") + + if args.generate_skills and args.disable_skills: + parser.error("Cannot use --generate-skills together with --disable-skills") + + if args.SCM == "filesystem" and args.disable_tools and not args.deep: + parser.error("Filesystem scans without --deep require tool use. Remove --disable-tools or add --deep.") + + if args.iterations < 1: + parser.error("--iterations must be at least 1") + + if args.iterations > 1 and (args.SCM != "filesystem" or args.deep): + sys.stdout.write( + f" โš ๏ธ warning: --iterations only applies to filesystem scans without --deep; ignoring --iterations={args.iterations}.{linesep}" + ) + + return args diff --git a/saist/util/caching.py b/saist/util/caching.py index 7e02df5..265bfde 100644 --- a/saist/util/caching.py +++ b/saist/util/caching.py @@ -7,6 +7,24 @@ async def hash_file(scm: Scm, filename: str) -> str: file: str = await scm.read_file_contents(filename) return hashlib.sha256(file.encode()).hexdigest() +async def hash_files(scm: Scm, filenames: list[str], extra: str = "") -> str: + hasher = hashlib.sha256() + if extra: + hasher.update(extra.encode("utf-8")) + hasher.update(b"\0") + + for filename in sorted(filenames): + hasher.update(filename.encode("utf-8")) + hasher.update(b"\0") + file_contents = await scm.read_file_contents(filename) + if file_contents is None: + hasher.update(b"") + else: + hasher.update(file_contents.encode("utf-8")) + hasher.update(b"\0") + + return hasher.hexdigest() + def finding_from_json_cache(json_dict: dict[str, any]) -> Finding: return Finding.model_validate(json_dict) @@ -23,4 +41,27 @@ def store_findings_to_cache_file(filename: str, findings: list[Finding], cache_f "findings": findings, } with open(cache_file, "w", encoding="utf-8") as cf: - json.dump(cache_dict, cf, cls=FindingJSONEncoder) \ No newline at end of file + json.dump(cache_dict, cf, cls=FindingJSONEncoder) + +def filesystem_tool_findings_from_cache_file(cache_file: str) -> tuple[list[Finding], set[str]]: + with open(cache_file, "r", encoding="utf-8") as file: + cache_json = json.load(file) + findings = cache_json.get("findings") or [] + files_read = cache_json.get("files_read") or [] + return [finding_from_json_cache(json_dict) for json_dict in findings], set(files_read) + +def store_filesystem_tool_findings_to_cache_file( + iteration: int, + filenames: list[str], + findings: list[Finding], + files_read: set[str], + cache_file: str, +): + cache_dict = { + "path": f"filesystem-tool-iteration-{iteration}", + "files": sorted(filenames), + "files_read": sorted(files_read), + "findings": findings, + } + with open(cache_file, "w", encoding="utf-8") as cf: + json.dump(cache_dict, cf, cls=FindingJSONEncoder) diff --git a/saist/util/output.py b/saist/util/output.py index c0a231e..0222201 100644 --- a/saist/util/output.py +++ b/saist/util/output.py @@ -33,8 +33,10 @@ def write_csv(findings: Iterable[Finding], csv_path: str): fieldnames = list(Finding.model_json_schema()["properties"].keys()) with open(csv_path, "w", newline="") as fp: - writer = csv.DictWriter(fp, fieldnames=fieldnames) - writer.writeheader() - for finding in findings: - writer.writerow(finding.model_dump()) - print(f"Written files to {csv_path}") + writer = csv.DictWriter(fp, fieldnames=fieldnames) + writer.writeheader() + for finding in findings: + row = finding.model_dump() + row["validation_steps"] = json.dumps(row.get("validation_steps", [])) + writer.writerow(row) + print(f"Written files to {csv_path}") diff --git a/saist/util/prompts.py b/saist/util/prompts.py index e77877a..af2d24e 100644 --- a/saist/util/prompts.py +++ b/saist/util/prompts.py @@ -1,32 +1,67 @@ -class prompts(): +class prompts: SUMMARY_PRE = """ - You are a senior application security engineer. - Given the following list of findings (issue descriptions and recommendations) - Write an informative summary suitable, and include headings - Group similar issues, and prioritize by severity. - Do not use any markdown - Return only the summary, no other text - """ +You are a senior application security engineer. +Write an informative executive summary suitable for a security review report. +Group similar issues, prioritize by severity, and make the business risk clear. +Do not use markdown. +Return only the summary, no other text. +""" + SUMMARY_POST = """ - findings: - """ +Findings: +""" + DETECT_PRE = """ - You are a security reviewer analyzing a single file's diff from a Pull Request. - Look for issues in the OWASP top ten. Identify as many as you can. - Report multiple issues per line as seperate findings. - When you detect a vulnerability get the full file by retrieving its contents, use this for context. - You can also retrieve other files for context as needed. - Only report a vulnerability if exists in the original diff. - Do not report vulnerabilities that exist only in tool output - Provide a vulnerability priority between 1 and 9. 9 is most critical - Map each finding to a Common Weakness Enumeration ID (CWE). - """ - DETECT_POST = """" - Below is the diff for this single file. It starts with 'File: ' followed by the unified diff.\n" - """ +You are a senior application security engineer. +Your job is to find exploitable vulnerabilities, not to produce a best-practice checklist. +Only report a finding when the supplied code supports a realistic attack path, privilege abuse path, data exposure path, or integrity impact. +Prefer business logic flaws and classic vulnerability classes over style, maintainability, hardening, or generic defense-in-depth advice. + +High-value issues include: +- Broken access control, tenant isolation failures, IDOR, missing ownership checks, unsafe role transitions, workflow bypasses, and confused-deputy flows. +- Authentication/session flaws, token misuse, password reset/invitation/account recovery abuse, MFA bypasses, and unsafe trust in client-controlled identity. +- Injection flaws such as SQL/NoSQL/LDAP/OS/template expression injection, unsafe deserialization, SSRF, path traversal, file upload abuse, command execution, and XSS where the sink is reachable. +- Secret exposure, insecure cryptography, dangerous framework configuration, webhook/signature verification failures, race conditions, payment/order/state-machine abuse, and unsafe admin or background-job actions. + +Do not report: +- Missing tests, missing logging, missing rate limits, missing security headers, missing CSP, missing validation, or missing error handling unless the code shows a concrete exploit path and reachable impact. +- Hypothetical risks that require unknown routes, unknown permissions, or assumptions not supported by code. +- Issues that only exist in tool output or in generated examples. + +For every finding, be specific about the attacker-controlled input or actor, the vulnerable code path, the security boundary crossed, and the impact. +For every finding, include concrete validation steps that a human reviewer can follow to reproduce or confirm the issue. These steps should identify the relevant entrypoint, required actor or permissions, input or request to try, expected vulnerable behavior, and the safe evidence that confirms impact. +If the evidence is weak, inspect more code with tools. If it is still weak, do not report it. +Report multiple vulnerabilities on the same line as separate findings only when they are distinct exploit paths. +Use the available tools to retrieve full files and related files when that context is needed. +Only report issues that are supported by the supplied code and context. +Provide a vulnerability priority between 1 and 9. 9 is most critical. +Map each finding to a Common Weakness Enumeration ID (CWE). +""" + + DETECT_POST = """ +Input format: +File: + +""" + + def detect(self, scm_prompt: str = "") -> str: + return "\n\n".join( + part.strip() + for part in [self.DETECT_PRE, scm_prompt, self.DETECT_POST] + if part and part.strip() + ) + + def summary(self, scm_prompt: str = "") -> str: + return "\n\n".join( + part.strip() + for part in [self.SUMMARY_PRE, scm_prompt, self.SUMMARY_POST] + if part and part.strip() + ) + @property def SUMMARY(self): - return self.SUMMARY_PRE + self.SUMMARY_POST + return self.summary() + @property def DETECT(self): - return self.DETECT_PRE + self.DETECT_POST + return self.detect() diff --git a/saist/util/skills.py b/saist/util/skills.py new file mode 100644 index 0000000..ab68749 --- /dev/null +++ b/saist/util/skills.py @@ -0,0 +1,469 @@ +import hashlib +import logging +import os +import re +from dataclasses import dataclass +from pathlib import Path +from typing import Annotated + +from pydantic import BaseModel, Field + +logger = logging.getLogger(__name__) + +DEFAULT_SKILLS_PATH = ".saist/skills" +DEFAULT_SKILL_MAX_BYTES = 60000 +DEFAULT_SKILL_SAMPLE_FILES = 80 +DEFAULT_SKILL_SAMPLE_BYTES = 12000 + +SKILL_SPECS = [ + ( + "authorization-model.md", + "How access control decisions are made, including roles, permissions, policies, ownership checks, and tenancy boundaries.", + ), + ( + "framework-specific-concerns.md", + "Framework conventions and security footguns that matter when reviewing this application.", + ), + ( + "input-validation-and-trust-boundaries.md", + "Where untrusted input enters the application and how validation, parsing, escaping, and serialization are handled.", + ), +] + +EXCLUDED_DIR_NAMES = { + ".git", + ".hg", + ".svn", + ".mypy_cache", + ".pytest_cache", + ".ruff_cache", + ".tox", + ".venv", + "SAISTCache", + "__pycache__", + "bin", + "build", + "coverage", + "dist", + "node_modules", + "obj", + "target", + "vendor", +} + +TEXT_EXTENSIONS = { + ".cs", + ".css", + ".env", + ".go", + ".graphql", + ".h", + ".hpp", + ".html", + ".java", + ".js", + ".json", + ".jsx", + ".kt", + ".md", + ".php", + ".py", + ".rb", + ".rs", + ".scala", + ".sh", + ".sql", + ".swift", + ".toml", + ".ts", + ".tsx", + ".xml", + ".yaml", + ".yml", +} + +IMPORTANT_FILENAMES = { + ".env.example", + "app.py", + "application.yml", + "application.yaml", + "build.gradle", + "cargo.toml", + "composer.json", + "docker-compose.yml", + "dockerfile", + "gemfile", + "go.mod", + "main.py", + "manage.py", + "middleware.py", + "package.json", + "pom.xml", + "program.cs", + "pyproject.toml", + "requirements.txt", + "routes.rb", + "settings.py", + "startup.cs", + "urls.py", +} + +IMPORTANT_PATH_TERMS = { + "admin", + "api", + "auth", + "config", + "controller", + "guard", + "handler", + "identity", + "middleware", + "migration", + "model", + "permission", + "policy", + "route", + "schema", + "security", + "service", + "session", + "tenant", + "user", + "validation", + "webhook", +} + + +@dataclass(frozen=True) +class AnalysisSkill: + path: Path + content: str + truncated: bool = False + + +@dataclass(frozen=True) +class SkillGenerationResult: + skills_dir: Path + written: list[Path] + skipped: list[Path] + sampled_files: int + + +class GeneratedSkillFile(BaseModel): + filename: Annotated[ + str, + Field(description="Markdown filename for the generated skill file, such as authorization-model.md"), + ] + content: Annotated[ + str, + Field(description="Durable Markdown instructions for future SAIST security analysis runs"), + ] + + +class GeneratedSkillFiles(BaseModel): + skills: list[GeneratedSkillFile] + + +def project_root_from_args(args) -> Path: + if getattr(args, "SCM", None) in {"filesystem", "git"} and getattr(args, "path", None): + return Path(args.path).expanduser().resolve() + + return Path.cwd().resolve() + + +def resolve_skills_dir(project_root: Path, skills_path: str) -> Path: + configured_path = Path(skills_path).expanduser() + if configured_path.is_absolute(): + return configured_path + + return (project_root / configured_path).resolve() + + +def load_analysis_skills(skills_dir: Path, max_bytes: int = DEFAULT_SKILL_MAX_BYTES) -> list[AnalysisSkill]: + if max_bytes <= 0: + return [] + + if not skills_dir.exists(): + return [] + + if not skills_dir.is_dir(): + logger.warning("Skills path exists but is not a directory: %s", skills_dir) + return [] + + skills: list[AnalysisSkill] = [] + remaining = max_bytes + + for path in sorted(skills_dir.rglob("*.md")): + if not path.is_file(): + continue + + try: + content = path.read_text(encoding="utf-8") + except UnicodeDecodeError: + logger.warning("Skill file is not valid UTF-8, skipping: %s", path) + continue + + content = content.strip() + if not content: + continue + + truncated = False + encoded_length = len(content.encode("utf-8")) + if encoded_length > remaining: + content = content.encode("utf-8")[:remaining].decode("utf-8", errors="ignore").strip() + truncated = True + + if content: + skills.append(AnalysisSkill(path=path, content=content, truncated=truncated)) + + remaining -= min(encoded_length, remaining) + if remaining <= 0: + break + + return skills + + +def format_analysis_skills(skills: list[AnalysisSkill]) -> str: + if not skills: + return "" + + sections = [ + "Application analysis skills:", + "The following project-specific skill files are guidance for this review. Use them to understand routing, identity, authorization, framework conventions, trust boundaries, and security-sensitive flows. Treat them as context, not proof of a vulnerability. Use this context to validate exploitability and business impact, not to report generic best-practice advice.", + ] + + for skill in skills: + marker = " (truncated)" if skill.truncated else "" + sections.append(f"\n--- {skill.path.name}{marker} ---\n{skill.content}") + + return "\n".join(sections) + + +def skills_prompt_digest(skills_prompt: str) -> str: + return hashlib.sha256(skills_prompt.encode("utf-8")).hexdigest()[:16] + + +async def generate_skill_files( + llm, + project_root: Path, + skills_dir: Path, + max_files: int = DEFAULT_SKILL_SAMPLE_FILES, + max_file_bytes: int = DEFAULT_SKILL_SAMPLE_BYTES, + overwrite: bool = False, +) -> SkillGenerationResult: + repository_profile, sampled_files = _build_repository_profile( + project_root=project_root, + skills_dir=skills_dir, + max_files=max_files, + max_file_bytes=max_file_bytes, + ) + + existing_skills = load_analysis_skills(skills_dir) + + system_prompt = """ +You are an application security architecture analyst generating durable SAIST skill files. +Skill files teach future security scans how this application works. They are not one-off vulnerability findings. + +Rules: +- Generate concise Markdown files with instructions future reviews can use. +- Use the requested filenames where possible. +- Do not invent facts. If the sampled files do not prove something, say what is unknown and what should be inspected. +- Capture conventions, security boundaries, review heuristics, and framework-specific risks. +- Avoid secrets, credentials, and long code excerpts. +- Focus on how to analyze future diffs in this application. +""" + + user_prompt = f""" +Generate SAIST skill files for this application. + +Requested skill files: +{_format_skill_specs()} + +Existing skill files, if any, should be preserved in spirit and improved from the sampled application context: +{_format_existing_skills(existing_skills)} + +Application context: +{repository_profile} +""" + + generated = await llm.prompt_structured(system_prompt, user_prompt, GeneratedSkillFiles) + + skills_dir.mkdir(parents=True, exist_ok=True) + written: list[Path] = [] + skipped: list[Path] = [] + + for skill in generated.skills: + filename = _safe_skill_filename(skill.filename) + if not filename: + continue + + output_path = skills_dir / filename + content = skill.content.strip() + if not content: + continue + + if output_path.exists() and not overwrite: + skipped.append(output_path) + continue + + output_path.write_text(content + "\n", encoding="utf-8") + written.append(output_path) + + return SkillGenerationResult( + skills_dir=skills_dir, + written=written, + skipped=skipped, + sampled_files=sampled_files, + ) + + +def _format_skill_specs() -> str: + return "\n".join(f"- {filename}: {description}" for filename, description in SKILL_SPECS) + + +def _format_existing_skills(skills: list[AnalysisSkill]) -> str: + if not skills: + return "No existing skill files were found." + + return "\n\n".join(f"--- {skill.path.name} ---\n{skill.content}" for skill in skills) + + +def _build_repository_profile( + project_root: Path, + skills_dir: Path, + max_files: int, + max_file_bytes: int, +) -> tuple[str, int]: + if not project_root.exists() or not project_root.is_dir(): + raise FileNotFoundError(f"Project root does not exist or is not a directory: {project_root}") + + candidates = _rank_candidate_files(project_root, skills_dir) + selected = candidates[:max(0, max_files)] + + inventory = "\n".join(f"- {relative_path}" for _, relative_path, _ in candidates[:300]) + excerpts = [] + + for _, relative_path, absolute_path in selected: + content = _read_sample(absolute_path, max_file_bytes) + if not content: + continue + excerpts.append(f"--- {relative_path} ---\n{content}") + + profile = [ + f"Project root: {project_root}", + "Repository file inventory, ranked for security architecture discovery:", + inventory or "No candidate files were found.", + "Selected file excerpts:", + "\n\n".join(excerpts) or "No readable file excerpts were found.", + ] + + return "\n\n".join(profile), len(selected) + + +def _rank_candidate_files(project_root: Path, skills_dir: Path) -> list[tuple[int, str, Path]]: + candidates: list[tuple[int, str, Path]] = [] + + for current_root, dirnames, filenames in os.walk(project_root): + current_path = Path(current_root) + dirnames[:] = [ + dirname + for dirname in dirnames + if not _should_skip_directory(current_path / dirname, skills_dir) + ] + + for filename in filenames: + absolute_path = current_path / filename + if not absolute_path.is_file() or _is_under(absolute_path, skills_dir): + continue + + try: + relative_path = absolute_path.relative_to(project_root) + except ValueError: + continue + + score = _score_path(relative_path) + if score <= 0: + continue + + candidates.append((score, relative_path.as_posix(), absolute_path)) + + candidates.sort(key=lambda item: (-item[0], item[1])) + return candidates + + +def _should_skip_directory(path: Path, skills_dir: Path) -> bool: + name = path.name + if name in EXCLUDED_DIR_NAMES: + return True + + if name.startswith(".") and name not in {".github"}: + return True + + return _is_under(path, skills_dir) + + +def _score_path(relative_path: Path) -> int: + path_text = relative_path.as_posix().lower() + filename = relative_path.name.lower() + suffix = relative_path.suffix.lower() + + if suffix and suffix not in TEXT_EXTENSIONS: + return 0 + + score = 0 + if filename in IMPORTANT_FILENAMES: + score += 100 + + for term in IMPORTANT_PATH_TERMS: + if term in path_text: + score += 20 + + if suffix in TEXT_EXTENSIONS: + score += 10 + + score -= min(len(relative_path.parts), 10) + return score + + +def _read_sample(path: Path, max_bytes: int) -> str: + if max_bytes <= 0: + return "" + + try: + with path.open("rb") as file: + data = file.read(max_bytes + 1) + except OSError as e: + logger.debug("Unable to read sample file %s: %s", path, e) + return "" + + if b"\x00" in data: + return "" + + truncated = len(data) > max_bytes + text = data[:max_bytes].decode("utf-8", errors="replace").strip() + if truncated: + text += "\n[truncated]" + + return text + + +def _safe_skill_filename(filename: str) -> str: + cleaned = Path(filename).name.lower().strip() + cleaned = re.sub(r"[^a-z0-9._-]+", "-", cleaned) + cleaned = cleaned.strip(".-_") + + if not cleaned: + return "" + + if not cleaned.endswith(".md"): + cleaned += ".md" + + return cleaned + + +def _is_under(path: Path, parent: Path) -> bool: + try: + path.resolve().relative_to(parent.resolve()) + return True + except ValueError: + return False diff --git a/saist/web/template.html b/saist/web/template.html index d9afd92..85f48a9 100644 --- a/saist/web/template.html +++ b/saist/web/template.html @@ -1,517 +1,1177 @@ - - - - - - SAIST - Findings - - - - - - - - - - - - - -
- -
-

SAIST - Findings

-

You can save this page locally for later using "Right Click + Save As..."

-
- -
-
- -
- -
- - - - - - - + + + + + + SAIST - Findings + + + + +
+
+
+

Punk Security

+

SAIST Findings

+

AI generated code review findings with validation steps, affected code, and local triage status. Status changes are stored in this browser.

+
+
+
Total0
+
Open0
+
False positives0
+
+
+ +
+
+ + + + + +
+
+ +
+
+ +
+
+ + + 0 selected +
+ Tip: right-click a cell to filter, copy, or mark a row. +
+ +
+
+ +
+
+ + + + + + diff --git a/tests/test_argparsing.py b/tests/test_argparsing.py new file mode 100644 index 0000000..f9b2e52 --- /dev/null +++ b/tests/test_argparsing.py @@ -0,0 +1,371 @@ +import pytest + +from util import argparsing + + +def parse_with(monkeypatch, argv): + monkeypatch.setattr(argparsing.sys, "argv", ["saist"] + argv) + return argparsing.parse_args() + + +def test_parse_args_sets_defaults_for_filesystem(monkeypatch, tmp_path): + args = parse_with(monkeypatch, ["--llm", "faike", "filesystem", str(tmp_path)]) + + assert args.SCM == "filesystem" + assert args.path == str(tmp_path) + assert args.path_for_comparison is None + assert args.llm == "faike" + assert args.llm_api_key is None + assert args.llm_model is None + assert args.llm_rate_limit == 10 + assert args.iterations == 1 + assert args.thinking == "medium" + assert args.ollama_base_uri == "http://localhost:11434" + assert args.openai_base_uri is None + assert args.azure_openai_endpoint is None + assert args.azure_openai_api_version is None + assert args.interactive is False + assert args.disable_tools is False + assert args.deep is False + assert args.web is False + assert args.web_port == 8080 + assert args.web_host == "127.0.0.1" + assert args.ci is False + assert args.csv is False + assert args.csv_path == "results.csv" + assert args.pdf is False + assert args.pdf_filename == "report.pdf" + assert args.disable_caching is False + assert args.cache_folder == "SAISTCache" + assert args.skills_path == ".saist/skills" + assert args.disable_skills is False + assert args.generate_skills is False + assert args.overwrite_skills is False + assert args.skills_max_bytes == 60000 + assert args.skills_sample_files == 80 + assert args.skills_sample_bytes == 12000 + assert args.project_name is None + assert args.skip_line_length_check is False + assert args.max_line_length == 1000 + assert args.include is None + assert args.exclude is None + assert args.dry_run is False + assert args.verbose == 0 + + +def test_parse_args_accepts_all_global_scan_options(monkeypatch, tmp_path): + args = parse_with( + monkeypatch, + [ + "--llm", + "openai", + "--llm-api-key", + "test-key", + "--llm-model", + "gpt-test", + "--llm-rate-limit", + "3", + "--iterations", + "10", + "--thinking", + "high", + "--ollama-base-uri", + "http://ollama.example", + "--openai-base-uri", + "http://openai.example", + "--azure-openai-endpoint", + "https://example.openai.azure.com/openai/v1/", + "--azure-openai-api-version", + "preview", + "--interactive", + "--disable-tools", + "--deep", + "--web", + "--web-port", + "9999", + "--web-host", + "0.0.0.0", + "--ci", + "--csv", + "--csv-path", + str(tmp_path / "results.csv"), + "--pdf", + "--pdf-filename", + str(tmp_path / "report.pdf"), + "--disable-caching", + "--cache-folder", + str(tmp_path / "cache"), + "--skills-path", + str(tmp_path / "skills"), + "--disable-skills", + "--skills-max-bytes", + "321", + "--skills-sample-files", + "4", + "--skills-sample-bytes", + "500", + "--project-name", + "Project X", + "--skip-line-length-check", + "--max-line-length", + "222", + "--include", + "**/*.py", + "-i", + "**/*.js", + "--exclude", + "build/", + "-e", + "*.min.js", + "--dry-run", + "-vv", + "filesystem", + str(tmp_path / "app"), + "--path-for-comparison", + str(tmp_path / "base"), + ], + ) + + assert args.SCM == "filesystem" + assert args.path == str(tmp_path / "app") + assert args.path_for_comparison == str(tmp_path / "base") + assert args.llm == "openai" + assert args.llm_api_key == "test-key" + assert args.llm_model == "gpt-test" + assert args.llm_rate_limit == 3 + assert args.iterations == 10 + assert args.thinking == "high" + assert args.ollama_base_uri == "http://ollama.example" + assert args.openai_base_uri == "http://openai.example" + assert args.azure_openai_endpoint == "https://example.openai.azure.com/openai/v1/" + assert args.azure_openai_api_version == "preview" + assert args.interactive is True + assert args.disable_tools is True + assert args.deep is True + assert args.web is True + assert args.web_port == 9999 + assert args.web_host == "0.0.0.0" + assert args.ci is True + assert args.csv is True + assert args.csv_path == str(tmp_path / "results.csv") + assert args.pdf is True + assert args.pdf_filename == str(tmp_path / "report.pdf") + assert args.disable_caching is True + assert args.cache_folder == str(tmp_path / "cache") + assert args.skills_path == str(tmp_path / "skills") + assert args.disable_skills is True + assert args.generate_skills is False + assert args.overwrite_skills is False + assert args.skills_max_bytes == 321 + assert args.skills_sample_files == 4 + assert args.skills_sample_bytes == 500 + assert args.project_name == "Project X" + assert args.skip_line_length_check is True + assert args.max_line_length == 222 + assert args.include == [["**/*.py"], ["**/*.js"]] + assert args.exclude == [["build/"], ["*.min.js"]] + assert args.dry_run is True + assert args.verbose == 2 + + +def test_parse_args_accepts_skill_generation_options(monkeypatch, tmp_path): + args = parse_with( + monkeypatch, + [ + "--llm", + "faike", + "--generate-skills", + "--overwrite-skills", + "--skills-path", + "custom-skills", + "--skills-max-bytes", + "123", + "--skills-sample-files", + "2", + "--skills-sample-bytes", + "3", + "filesystem", + str(tmp_path), + ], + ) + + assert args.llm == "faike" + assert args.SCM == "filesystem" + assert args.path == str(tmp_path) + assert args.generate_skills is True + assert args.overwrite_skills is True + assert args.skills_path == "custom-skills" + assert args.skills_max_bytes == 123 + assert args.skills_sample_files == 2 + assert args.skills_sample_bytes == 3 + + +def test_parse_args_accepts_disabled_thinking(monkeypatch, tmp_path): + args = parse_with(monkeypatch, ["--llm", "faike", "--thinking", "disabled", "filesystem", str(tmp_path)]) + + assert args.thinking == "disabled" + + +def test_parse_args_accepts_open_ai_baseuri_alias(monkeypatch, tmp_path): + args = parse_with( + monkeypatch, + [ + "--llm", + "openai", + "--llm-api-key", + "test-key", + "--open-ai-baseuri", + "http://openai.example/v1", + "filesystem", + str(tmp_path), + ], + ) + + assert args.openai_base_uri == "http://openai.example/v1" + + +def test_parse_args_accepts_azure_foundry_without_generic_api_key(monkeypatch, tmp_path): + monkeypatch.setenv("AZURE_OPENAI_ENDPOINT", "https://example.openai.azure.com/openai/v1/") + monkeypatch.setenv("AZURE_OPENAI_API_KEY", "azure-key") + + args = parse_with(monkeypatch, ["--llm", "azure-foundry", "--llm-model", "gpt-5-mini", "filesystem", str(tmp_path)]) + + assert args.llm == "azure-foundry" + assert args.llm_model == "gpt-5-mini" + assert args.llm_api_key is None + assert args.azure_openai_endpoint is None + + +def test_parse_args_rejects_zero_iterations(monkeypatch, tmp_path): + with pytest.raises(SystemExit): + parse_with(monkeypatch, ["--llm", "faike", "--iterations", "0", "filesystem", str(tmp_path)]) + + +@pytest.mark.parametrize( + "argv", + [ + ["--llm", "faike", "--iterations", "3", "--deep", "filesystem", "/tmp/project"], + ["--llm", "faike", "--iterations", "3", "git", "/tmp/project"], + ["--llm", "faike", "--iterations", "3", "github", "owner/repo", "--github-token", "token", "123"], + ], +) +def test_parse_args_warns_when_iterations_are_ignored(monkeypatch, capsys, argv): + args = parse_with(monkeypatch, argv) + + assert args.iterations == 3 + assert "--iterations only applies to filesystem scans without --deep" in capsys.readouterr().out + + +def test_parse_args_accepts_git_subcommand_options(monkeypatch, tmp_path): + args = parse_with( + monkeypatch, + [ + "--llm", + "faike", + "git", + str(tmp_path), + "--ref-for-compare", + "develop", + "--ref-to-compare", + "feature", + "--commit-for-compare", + "abc123", + "--commit-to-compare", + "def456", + ], + ) + + assert args.SCM == "git" + assert args.path == str(tmp_path) + assert args.ref_for_compare == "develop" + assert args.ref_to_compare == "feature" + assert args.commit_for_compare == "abc123" + assert args.commit_to_compare == "def456" + + +def test_parse_args_sets_git_subcommand_defaults(monkeypatch, tmp_path): + args = parse_with(monkeypatch, ["--llm", "faike", "git", str(tmp_path)]) + + assert args.SCM == "git" + assert args.path == str(tmp_path) + assert args.ref_for_compare == "main" + assert args.ref_to_compare == "HEAD" + assert args.commit_for_compare is None + assert args.commit_to_compare is None + + +def test_parse_args_accepts_github_subcommand_options(monkeypatch): + args = parse_with( + monkeypatch, + [ + "--llm", + "faike", + "github", + "owner/repo", + "--github-token", + "github-token", + "123", + ], + ) + + assert args.SCM == "github" + assert args.repository == "owner/repo" + assert args.github_token == "github-token" + assert args.pr == "123" + + +def test_parse_args_accepts_poem_subcommand(monkeypatch): + args = parse_with(monkeypatch, ["--llm", "faike", "poem"]) + + assert args.SCM == "poem" + + +@pytest.mark.parametrize( + ("argv", "message"), + [ + ( + ["--llm", "openai", "filesystem", "/tmp/project"], + "You must provide an api key", + ), + ( + ["--llm", "faike", "--interactive", "filesystem", "/tmp/project"], + "Faike LLM", + ), + ( + ["--llm", "ollama", "--interactive", "filesystem", "/tmp/project"], + "cannot use the interactive shell with ollama", + ), + ( + ["--llm", "faike", "--generate-skills", "poem"], + "Cannot generate SAIST skills while using the poem command", + ), + ( + ["--llm", "faike", "--generate-skills", "--disable-skills", "filesystem", "/tmp/project"], + "Cannot use --generate-skills together with --disable-skills", + ), + ( + ["--llm", "faike", "--disable-tools", "filesystem", "/tmp/project"], + "Filesystem scans without --deep require tool use", + ), + ], +) +def test_parse_args_rejects_invalid_combinations(monkeypatch, capsys, argv, message): + with pytest.raises(SystemExit) as exc_info: + parse_with(monkeypatch, argv) + + assert exc_info.value.code == 2 + assert message in capsys.readouterr().out + + +def test_parse_args_rejects_bedrock_api_key(monkeypatch, capsys): + with pytest.raises(SystemExit) as exc_info: + parse_with(monkeypatch, ["--llm", "bedrock", "--llm-api-key", "nope", "filesystem", "/tmp/project"]) + + assert exc_info.value.code == 2 + assert "Do not provide an API key for bedrock" in capsys.readouterr().out + + +def test_parse_args_accepts_pdf_without_external_renderer_dependency(monkeypatch): + args = parse_with(monkeypatch, ["--llm", "faike", "--pdf", "filesystem", "/tmp/project"]) + + assert args.pdf is True diff --git a/tests/test_caching_output.py b/tests/test_caching_output.py new file mode 100644 index 0000000..cd1f2d7 --- /dev/null +++ b/tests/test_caching_output.py @@ -0,0 +1,111 @@ +import asyncio +import csv +import json + +from models import Finding +from util.caching import ( + filesystem_tool_findings_from_cache_file, + findings_from_cache_file, + hash_file, + hash_files, + store_filesystem_tool_findings_to_cache_file, + store_findings_to_cache_file, +) +from util.output import write_csv + + +def make_finding(**overrides): + values = { + "file": "app.py", + "snippet": "danger()", + "title": "Dangerous call", + "issue": "A dangerous call was introduced.", + "recommendation": "Remove the dangerous call.", + "validation_steps": ["Call the affected endpoint.", "Confirm the dangerous behavior is reachable."], + "cwe": "CWE-20", + "priority": 5, + "line_number": 7, + } + values.update(overrides) + return Finding.model_validate(values) + + +def test_hash_file_hashes_scm_file_contents(): + class FakeScm: + async def read_file_contents(self, filename): + assert filename == "app.py" + return "same content\n" + + assert asyncio.run(hash_file(FakeScm(), "app.py")) == "f953bbd204bb867e48a6ff774cffa3dcffd02c6580e8f1d00c37dbbaa743d6c8" + + +def test_hash_files_includes_file_names_contents_and_extra_context(): + class FakeScm: + async def read_file_contents(self, filename): + return {"app.py": "same content\n", "binary.gz": None}[filename] + + first = asyncio.run(hash_files(FakeScm(), ["binary.gz", "app.py"], extra="iteration prompt")) + second = asyncio.run(hash_files(FakeScm(), ["app.py", "binary.gz"], extra="iteration prompt")) + different_extra = asyncio.run(hash_files(FakeScm(), ["app.py", "binary.gz"], extra="other prompt")) + + assert first == second + assert first != different_extra + + +def test_findings_cache_round_trip(tmp_path): + cache_file = tmp_path / "finding.json" + finding = make_finding(priority=8) + + store_findings_to_cache_file("app.py", [finding], str(cache_file)) + loaded = findings_from_cache_file(str(cache_file)) + + assert loaded == [finding] + + +def test_filesystem_tool_findings_cache_round_trip(tmp_path): + cache_file = tmp_path / "shallow.json" + finding = make_finding(priority=8) + + store_filesystem_tool_findings_to_cache_file( + iteration=2, + filenames=["app.py", "settings.py"], + findings=[finding], + files_read={"settings.py"}, + cache_file=str(cache_file), + ) + + loaded_findings, files_read = filesystem_tool_findings_from_cache_file(str(cache_file)) + + assert loaded_findings == [finding] + assert files_read == {"settings.py"} + + +def test_findings_cache_returns_empty_list_for_null_findings(tmp_path): + cache_file = tmp_path / "empty.json" + cache_file.write_text(json.dumps({"path": "app.py", "findings": None}), encoding="utf-8") + + assert findings_from_cache_file(str(cache_file)) == [] + + +def test_write_csv_writes_finding_fields(tmp_path): + csv_path = tmp_path / "findings.csv" + finding = make_finding(cwe="CWE-89", priority=9) + + write_csv([finding], str(csv_path)) + + with csv_path.open(newline="") as file: + rows = list(csv.DictReader(file)) + + assert rows == [ + { + "file": "app.py", + "snippet": "danger()", + "title": "Dangerous call", + "issue": "A dangerous call was introduced.", + "recommendation": "Remove the dangerous call.", + "validation_steps": '["Call the affected endpoint.", "Confirm the dangerous behavior is reachable."]', + "cwe": "CWE-89", + "priority": "9", + "line_number": "7", + } + ] diff --git a/tests/test_filesystem_adapter.py b/tests/test_filesystem_adapter.py new file mode 100644 index 0000000..fdedaaa --- /dev/null +++ b/tests/test_filesystem_adapter.py @@ -0,0 +1,107 @@ +import asyncio +import logging + +from scm.adapters.filesystem import FilesystemAdapter + + +def test_filesystem_adapter_lists_text_files_and_skips_binary_files(tmp_path): + (tmp_path / "app.py").write_text("print('hello')\n", encoding="utf-8") + (tmp_path / "image.bin").write_bytes(b"\xff\xfe\x00\x00") + + adapter = FilesystemAdapter(compare_path=str(tmp_path)) + changed_files = adapter.get_changed_files() + + assert [file["filename"] for file in changed_files] == ["app.py"] + assert "+print('hello')" in changed_files[0]["patch"] + + +def test_filesystem_adapter_reads_file_contents(tmp_path): + (tmp_path / "app.py").write_text("print('hello')\n", encoding="utf-8") + adapter = FilesystemAdapter(compare_path=str(tmp_path)) + + contents = asyncio.run(adapter.get_file_contents("app.py")) + + assert contents == "print('hello')\n" + + +def test_filesystem_adapter_returns_none_for_gzip_or_binary_file_without_err_log(tmp_path, caplog): + (tmp_path / "archive.gz").write_bytes(b"\x1f\x8b\x08\x00not text") + adapter = FilesystemAdapter(compare_path=str(tmp_path)) + + with caplog.at_level(logging.WARNING): + contents = asyncio.run(adapter.get_file_contents("archive.gz")) + + assert contents is None + assert "ERR:" not in caplog.text + assert "codec can't decode" not in caplog.text + + +def test_filesystem_adapter_lists_all_files(tmp_path): + (tmp_path / "app.py").write_text("print('hello')\n", encoding="utf-8") + (tmp_path / "image.bin").write_bytes(b"\xff\xfe\x00\x00") + (tmp_path / "nested").mkdir() + (tmp_path / "nested" / "settings.py").write_text("DEBUG = True\n", encoding="utf-8") + + adapter = FilesystemAdapter(compare_path=str(tmp_path)) + + assert asyncio.run(adapter.list_files()) == ["app.py", "image.bin", "nested/settings.py"] + + +def test_filesystem_adapter_regex_searches_text_files(tmp_path): + (tmp_path / "app.py").write_text("SECRET_KEY = 'dev'\nprint(SECRET_KEY)\n", encoding="utf-8") + (tmp_path / "README.md").write_text("SECRET_KEY is documented here\n", encoding="utf-8") + (tmp_path / "image.bin").write_bytes(b"\xff\xfe\x00\x00") + (tmp_path / "archive.gz").write_bytes(b"\x1f\x8b\x08\x00SECRET_KEY") + + adapter = FilesystemAdapter(compare_path=str(tmp_path)) + matches = asyncio.run(adapter.regex_search(r"SECRET_KEY", file_pattern="**/*.py")) + + assert matches == [ + { + "filename": "app.py", + "line_number": 1, + "column": 1, + "match": "SECRET_KEY", + "line": "SECRET_KEY = 'dev'", + }, + { + "filename": "app.py", + "line_number": 2, + "column": 7, + "match": "SECRET_KEY", + "line": "print(SECRET_KEY)", + }, + ] + + +def test_filesystem_adapter_regex_search_returns_error_for_invalid_regex(tmp_path): + adapter = FilesystemAdapter(compare_path=str(tmp_path)) + + result = asyncio.run(adapter.regex_search("[")) + + assert result[0]["error"].startswith("Invalid regex:") + assert result[0]["pattern"] == "[" + + +def test_filesystem_adapter_returns_none_for_missing_or_outside_file(tmp_path): + adapter = FilesystemAdapter(compare_path=str(tmp_path)) + + assert asyncio.run(adapter.get_file_contents("missing.py")) is None + assert asyncio.run(adapter.get_file_contents("/etc/passwd")) is None + + +def test_filesystem_adapter_compares_base_and_compare_paths(tmp_path): + base_path = tmp_path / "base" + compare_path = tmp_path / "compare" + base_path.mkdir() + compare_path.mkdir() + (base_path / "app.py").write_text("value = 'old'\n", encoding="utf-8") + (compare_path / "app.py").write_text("value = 'new'\n", encoding="utf-8") + (compare_path / "new_only.py").write_text("ignored by current dircmp implementation\n", encoding="utf-8") + + adapter = FilesystemAdapter(compare_path=str(compare_path), base_path=str(base_path)) + changed_files = adapter.get_changed_files() + + assert [file["filename"] for file in changed_files] == ["app.py"] + assert "-value = 'old'" in changed_files[0]["patch"] + assert "+value = 'new'" in changed_files[0]["patch"] diff --git a/tests/test_filtering.py b/tests/test_filtering.py new file mode 100644 index 0000000..372fe01 --- /dev/null +++ b/tests/test_filtering.py @@ -0,0 +1,46 @@ +from util.filtering import FilterRules + + +def test_filter_rules_apply_include_and_exclude_files(tmp_path): + include_file = tmp_path / "saist.include" + exclude_file = tmp_path / "saist.ignore" + include_file.write_text("src/**/*.py\nREADME.md\n", encoding="utf-8") + exclude_file.write_text("src/generated/\n", encoding="utf-8") + + rules = FilterRules( + include_patterns=None, + exclude_patterns=None, + include_rules_file=include_file, + exclude_rules_file=exclude_file, + ) + + assert rules.filename_included(str(tmp_path / "src" / "app.py")) + assert rules.filename_included(str(tmp_path / "README.md")) + assert not rules.filename_included(str(tmp_path / "src" / "generated" / "client.py")) + assert not rules.filename_included(str(tmp_path / "src" / "app.js")) + + +def test_filter_rules_allow_cli_patterns_to_extend_file_rules(tmp_path): + include_file = tmp_path / "saist.include" + exclude_file = tmp_path / "saist.ignore" + include_file.write_text("src/**/*.py\n", encoding="utf-8") + exclude_file.write_text("", encoding="utf-8") + + rules = FilterRules( + include_patterns=[["tools/**/*.sh"]], + exclude_patterns=[["**/danger.sh"]], + include_rules_file=include_file, + exclude_rules_file=exclude_file, + ) + + assert rules.filename_included(str(tmp_path / "src" / "app.py")) + assert rules.filename_included(str(tmp_path / "tools" / "run.sh")) + assert not rules.filename_included(str(tmp_path / "tools" / "danger.sh")) + + +def test_file_exceeds_line_length_limit_checks_file_and_patch_text(): + rules = FilterRules(include_patterns=None, exclude_patterns=None) + + assert rules.file_exceeds_line_length_limit("short\n", "+short\n", max_line_length=10) is False + assert rules.file_exceeds_line_length_limit("x" * 11, "+short\n", max_line_length=10) is True + assert rules.file_exceeds_line_length_limit("short\n", "+" + "x" * 11, max_line_length=10) is True diff --git a/tests/test_git.py b/tests/test_git.py new file mode 100644 index 0000000..712001e --- /dev/null +++ b/tests/test_git.py @@ -0,0 +1,40 @@ +from util.git import parse_unified_diff + + +def test_parse_unified_diff_maps_added_and_context_lines(): + patch = """diff --git a/app.py b/app.py +index 0000000..1111111 100644 +--- a/app.py ++++ b/app.py +@@ -1,3 +1,4 @@ + import os +-query = "safe" ++query = request.args["q"] ++db.execute("select * from users where name = " + query) + print(query) +""" + + line_map, new_lines_text = parse_unified_diff(patch) + + assert new_lines_text == { + 1: "import os", + 2: 'query = request.args["q"]', + 3: 'db.execute("select * from users where name = " + query)', + 4: "print(query)", + } + assert set(line_map) == {1, 2, 3, 4} + assert line_map[2] < line_map[3] + + +def test_parse_unified_diff_ignores_headers_before_first_hunk(): + line_map, new_lines_text = parse_unified_diff( + """diff --git a/app.py b/app.py +metadata that should not be parsed +@@ -0,0 +1,2 @@ ++first = True ++second = True +""" + ) + + assert new_lines_text == {1: "first = True", 2: "second = True"} + assert set(line_map) == {1, 2} diff --git a/tests/test_git_adapter.py b/tests/test_git_adapter.py new file mode 100644 index 0000000..cd417cf --- /dev/null +++ b/tests/test_git_adapter.py @@ -0,0 +1,68 @@ +import asyncio + +from git import Actor, Repo + +from scm.adapters.git import GitAdapter + + +AUTHOR = Actor("SAIST Tests", "tests@example.com") + + +def commit(repo, message): + return repo.index.commit(message, author=AUTHOR, committer=AUTHOR) + + +def test_git_adapter_reads_and_lists_files_at_compare_commit(tmp_path): + repo = Repo.init(tmp_path) + (tmp_path / "app.py").write_text("print('hello')\n", encoding="utf-8") + (tmp_path / "nested").mkdir() + (tmp_path / "nested" / "settings.py").write_text("DEBUG = True\n", encoding="utf-8") + (tmp_path / "image.bin").write_bytes(b"\xff\xfe\x00\x00") + repo.index.add(["app.py", "nested/settings.py", "image.bin"]) + head = commit(repo, "initial") + + adapter = GitAdapter(repo_path=tmp_path, base_commit=head.hexsha, compare_commit=head.hexsha) + + assert asyncio.run(adapter.list_files()) == ["app.py", "image.bin", "nested/settings.py"] + assert asyncio.run(adapter.get_file_contents("nested/settings.py")) == "DEBUG = True\n" + assert asyncio.run(adapter.get_file_contents("../outside.py")) is None + assert asyncio.run(adapter.get_file_contents("image.bin")) is None + + +def test_git_adapter_regex_searches_compare_commit_files(tmp_path): + repo = Repo.init(tmp_path) + (tmp_path / "app.py").write_text("SECRET_KEY = 'dev'\nprint(SECRET_KEY)\n", encoding="utf-8") + (tmp_path / "README.md").write_text("SECRET_KEY is documented here\n", encoding="utf-8") + repo.index.add(["app.py", "README.md"]) + head = commit(repo, "initial") + + adapter = GitAdapter(repo_path=tmp_path, base_commit=head.hexsha, compare_commit=head.hexsha) + matches = asyncio.run(adapter.regex_search(r"SECRET_KEY", file_pattern="**/*.py", max_results=1)) + + assert matches == [ + { + "filename": "app.py", + "line_number": 1, + "column": 1, + "match": "SECRET_KEY", + "line": "SECRET_KEY = 'dev'", + } + ] + + +def test_git_adapter_get_changed_files_uses_base_to_compare_patch(tmp_path): + repo = Repo.init(tmp_path) + (tmp_path / "app.py").write_text("value = 'old'\n", encoding="utf-8") + repo.index.add(["app.py"]) + base = commit(repo, "base") + + (tmp_path / "app.py").write_text("value = 'new'\n", encoding="utf-8") + repo.index.add(["app.py"]) + compare = commit(repo, "compare") + + adapter = GitAdapter(repo_path=tmp_path, base_commit=base.hexsha, compare_commit=compare.hexsha) + changed_files = adapter.get_changed_files() + + assert [file["filename"] for file in changed_files] == ["app.py"] + assert "-value = 'old'" in changed_files[0]["patch"] + assert "+value = 'new'" in changed_files[0]["patch"] diff --git a/tests/test_github_adapter.py b/tests/test_github_adapter.py new file mode 100644 index 0000000..7f33cb7 --- /dev/null +++ b/tests/test_github_adapter.py @@ -0,0 +1,72 @@ +import asyncio +import base64 + +from scm.adapters.github import Github + + +class FakeResponse: + def __init__(self, payload=None, status_code=200): + self.payload = payload + self.status_code = status_code + + def json(self): + return self.payload + + def raise_for_status(self): + if self.status_code >= 400: + import requests + + raise requests.exceptions.HTTPError(response=self) + + +class FakeSession: + def __init__(self, responses): + self.responses = list(responses) + self.urls = [] + + def get(self, url, **kwargs): + self.urls.append(url) + return self.responses.pop(0) + + +def github_with_session(session): + adapter = Github.__new__(Github) + adapter.requests = session + adapter.repo = "owner/repo" + adapter.pr_number = "123" + adapter.commit_sha = "abc123" + return adapter + + +def test_github_changed_files_excludes_removed_files(): + session = FakeSession( + [ + FakeResponse( + [ + {"filename": "app.py", "status": "modified", "patch": "@@ -1 +1 @@"}, + {"filename": "deleted.py", "status": "removed", "patch": "@@ -1 +0 @@"}, + ] + ), + FakeResponse([]), + ] + ) + adapter = github_with_session(session) + + changed_files = adapter.get_changed_files() + + assert changed_files == [{"filename": "app.py", "status": "modified", "patch": "@@ -1 +1 @@"}] + + +def test_github_get_file_contents_returns_none_for_missing_pr_head_file(): + session = FakeSession([FakeResponse(status_code=404)]) + adapter = github_with_session(session) + + assert asyncio.run(adapter.get_file_contents("deleted.py")) is None + + +def test_github_get_file_contents_decodes_base64_content(): + content = base64.b64encode(b"print('hello')\n").decode() + session = FakeSession([FakeResponse({"encoding": "base64", "content": content})]) + adapter = github_with_session(session) + + assert asyncio.run(adapter.get_file_contents("app.py")) == "print('hello')\n" diff --git a/tests/test_llm_adapters.py b/tests/test_llm_adapters.py new file mode 100644 index 0000000..ec048d0 --- /dev/null +++ b/tests/test_llm_adapters.py @@ -0,0 +1,94 @@ +import asyncio + +from llm import adapters +from llm.adapters import BaseLlmAdapter +from llm.adapters.azure_foundry import AzureFoundryAdapter +from llm.adapters.faike import FaikeAdapter +from llm.adapters.openai import OpenAiAdapter +from models import Findings +from pydantic_ai.models.openai import OpenAIResponsesModel + + +def test_model_options_include_pydantic_thinking_level_when_supported(monkeypatch): + monkeypatch.setattr(adapters, "pydantic_ai_supports_thinking", lambda: True) + adapter = BaseLlmAdapter(thinking="high") + + assert adapter.get_model_options()["thinking"] == "high" + assert "temperature" not in adapter.get_model_options() + + +def test_model_options_drop_custom_sampling_parameters_when_thinking_enabled(monkeypatch): + monkeypatch.setattr(adapters, "pydantic_ai_supports_thinking", lambda: True) + adapter = BaseLlmAdapter(thinking="medium") + adapter.model_options = {"temperature": 0.7, "timeout": 30} + + options = adapter.get_model_options() + + assert options["thinking"] == "medium" + assert options["timeout"] == 30 + assert "temperature" not in options + + +def test_model_options_map_disabled_thinking_to_false(monkeypatch): + monkeypatch.setattr(adapters, "pydantic_ai_supports_thinking", lambda: True) + adapter = BaseLlmAdapter(thinking="disabled") + + options = adapter.get_model_options() + + assert options["thinking"] is False + assert options["temperature"] == 0.0 + + +def test_model_options_omit_thinking_when_pydantic_ai_does_not_support_it(monkeypatch): + monkeypatch.setattr(adapters, "pydantic_ai_supports_thinking", lambda: False) + adapter = BaseLlmAdapter(thinking="xhigh") + + assert "thinking" not in adapter.get_model_options() + + +def test_faike_adapter_accepts_thinking_without_using_it(): + adapter = FaikeAdapter("", "Fake LLM", thinking="low") + + result = asyncio.run( + adapter.prompt_structured( + "system", + "File: app.py\n@@ -0,0 +1 @@\n+[]\n", + Findings, + ) + ) + + assert adapter.thinking == "low" + assert result.findings[0].file == "app.py" + + +def test_openai_adapter_uses_responses_model_for_thinking_support(): + adapter = OpenAiAdapter(model="gpt-5-mini", api_key="test-key", thinking="high") + + assert isinstance(adapter.model, OpenAIResponsesModel) + assert adapter.model_name == "gpt-5-mini" + assert adapter.get_model_options()["thinking"] == "high" + + +def test_openai_adapter_uses_custom_base_url(): + adapter = OpenAiAdapter( + model="gpt-5-mini", + api_key="test-key", + base_url="http://openai.example/v1", + ) + + assert str(adapter.model.provider.client.base_url) == "http://openai.example/v1/" + + +def test_azure_foundry_adapter_uses_responses_model_and_azure_provider(): + adapter = AzureFoundryAdapter( + model="gpt-5-mini", + api_key="test-key", + azure_endpoint="https://example.openai.azure.com/openai/v1/", + thinking="high", + ) + + assert isinstance(adapter.model, OpenAIResponsesModel) + assert adapter.model_name == "gpt-5-mini" + assert adapter.model_vendor == "Azure AI Foundry" + assert adapter.model.provider.name == "azure" + assert adapter.get_model_options()["thinking"] == "high" diff --git a/tests/test_main_analysis.py b/tests/test_main_analysis.py new file mode 100644 index 0000000..962e27b --- /dev/null +++ b/tests/test_main_analysis.py @@ -0,0 +1,847 @@ +import asyncio +import time +from types import SimpleNamespace + +import main as saist_main +from models import Finding, Findings +from scm.adapters.filesystem import FilesystemAdapter +from scm.adapters.git import GitAdapter as RealGitAdapter +from scm.adapters.github import Github as RealGithub +from util.skills import skills_prompt_digest + + +def test_analyze_single_file_includes_analysis_skills_in_system_prompt(): + class CapturingLlm: + def __init__(self): + self.system_prompt = None + self.user_prompt = None + self.tool_fns = None + + async def prompt_structured(self, system_prompt, user_prompt, response_format, tool_fns=None): + self.system_prompt = system_prompt + self.user_prompt = user_prompt + self.tool_fns = tool_fns + return Findings(findings=[]) + + class FakeScm: + async def read_file_contents(self, filename): + return "print('hello')\n" + + async def list_files(self): + return ["app.py"] + + async def regex_search(self, pattern, file_pattern="**/*", max_results=100): + return [] + + def tool_functions(self): + return [self.read_file_contents, self.list_files, self.regex_search] + + def detect_prompt(self): + return FilesystemAdapter.DETECT_PROMPT + + llm = CapturingLlm() + result = asyncio.run( + saist_main.analyze_single_file( + scm=FakeScm(), + adapter=llm, + filename="app.py", + patch_text="@@ -0,0 +1 @@\n+print('hello')\n", + disable_tools=False, + analysis_skills="Skill: authorization requires tenant ownership checks.", + ) + ) + + assert result == [] + assert "Skill: authorization requires tenant ownership checks." in llm.system_prompt + assert "File: app.py" in llm.user_prompt + assert [tool.__name__ for tool in llm.tool_fns] == [ + "read_file_contents", + "list_files", + "regex_search", + ] + assert "not to produce a best-practice checklist" in llm.system_prompt + assert "penetration test style review across the entire application codebase" in llm.system_prompt + assert "cross-tenant data access" in llm.system_prompt + assert "include concrete validation steps" in llm.system_prompt + + +def test_analyze_single_file_uses_git_diff_prompt_for_git_adapter(): + class CapturingLlm: + def __init__(self): + self.system_prompt = None + + async def prompt_structured(self, system_prompt, user_prompt, response_format, tool_fns=None): + self.system_prompt = system_prompt + return Findings(findings=[]) + + class GitAdapter: + def tool_functions(self): + return [] + + def detect_prompt(self): + return RealGitAdapter.DETECT_PROMPT + + llm = CapturingLlm() + result = asyncio.run( + saist_main.analyze_single_file( + scm=GitAdapter(), + adapter=llm, + filename="app.py", + patch_text="@@ -1 +1 @@\n-old\n+new\n", + disable_tools=False, + ) + ) + + assert result == [] + assert "analyzing a diff of code that needs security review" in llm.system_prompt + assert "git comparison" in llm.system_prompt + assert "anchored to changed lines" in llm.system_prompt + assert "pre-existing best-practice issues" in llm.system_prompt + assert "penetration test style review" not in llm.system_prompt + + +def test_analyze_single_file_uses_github_pull_request_prompt(): + class CapturingLlm: + def __init__(self): + self.system_prompt = None + + async def prompt_structured(self, system_prompt, user_prompt, response_format, tool_fns=None): + self.system_prompt = system_prompt + return Findings(findings=[]) + + class Github: + def tool_functions(self): + return [] + + def detect_prompt(self): + return RealGithub.DETECT_PROMPT + + llm = CapturingLlm() + result = asyncio.run( + saist_main.analyze_single_file( + scm=Github(), + adapter=llm, + filename="app.py", + patch_text="@@ -1 +1 @@\n-old\n+new\n", + disable_tools=False, + ) + ) + + assert result == [] + assert "GitHub pull request" in llm.system_prompt + assert "Report only vulnerabilities anchored to changed lines" in llm.system_prompt + assert "business logic changes" in llm.system_prompt + + +def test_filesystem_tool_analysis_sends_file_inventory_and_tracks_coverage(): + class CapturingLlm: + def __init__(self): + self.system_prompt = None + self.user_prompt = None + self.tool_names = None + + async def prompt_structured(self, system_prompt, user_prompt, response_format, tool_fns=None): + self.system_prompt = system_prompt + self.user_prompt = user_prompt + self.tool_names = [tool.__name__ for tool in tool_fns] + read_file_contents = next(tool for tool in tool_fns if tool.__name__ == "read_file_contents") + regex_search = next(tool for tool in tool_fns if tool.__name__ == "regex_search") + await read_file_contents("app.py") + await regex_search("SECRET", "**/*.py", 10) + return Findings( + findings=[ + Finding( + file="app.py", + snippet="SECRET", + title="Secret", + issue="Issue", + recommendation="Fix it.", + validation_steps=["Search for SECRET and confirm it is committed."], + cwe="CWE-798", + priority=6, + line_number=1, + ) + ] + ) + + class FakeScm: + async def read_file_contents(self, filename): + return "SECRET = 'dev'\n" + + async def list_files(self): + return ["app.py", "settings.py"] + + async def regex_search(self, pattern, file_pattern="**/*", max_results=100): + return [{"filename": "settings.py", "line_number": 1, "column": 1, "match": "SECRET", "line": "SECRET = 'dev'"}] + + def detect_prompt(self): + return FilesystemAdapter.DETECT_PROMPT + + llm = CapturingLlm() + findings, files_read = asyncio.run( + saist_main.generate_findings_with_filesystem_tools( + scm=FakeScm(), + llm=llm, + filenames=["app.py", "settings.py"], + disable_tools=False, + analysis_skills="Skill guidance", + ) + ) + + assert [finding.file for finding in findings] == ["app.py"] + assert files_read == {"app.py", "settings.py"} + assert "Application file inventory" in llm.user_prompt + assert "- app.py" in llm.user_prompt + assert "- settings.py" in llm.user_prompt + assert "penetration test style review across the entire application codebase" in llm.system_prompt + assert "Trace attacker-controlled input from entrypoint to sink" in llm.system_prompt + assert "Skill guidance" in llm.system_prompt + assert llm.tool_names == ["read_file_contents", "list_files", "regex_search"] + + +def test_filesystem_tool_analysis_iterations_respect_concurrency_limit(): + class CountingLlm: + def __init__(self): + self.calls = 0 + self.active = 0 + self.max_active = 0 + + async def prompt_structured(self, system_prompt, user_prompt, response_format, tool_fns=None): + self.calls += 1 + self.active += 1 + self.max_active = max(self.max_active, self.active) + await asyncio.sleep(0.01) + self.active -= 1 + return Findings( + findings=[ + Finding( + file="app.py", + snippet="SECRET", + title=f"Secret {self.calls}", + issue=f"Issue {self.calls}", + recommendation="Fix it.", + validation_steps=["Confirm the issue is reachable."], + cwe="CWE-798", + priority=6, + line_number=1, + ) + ] + ) + + class FakeScm: + async def read_file_contents(self, filename): + return "SECRET = 'dev'\n" + + async def list_files(self): + return ["app.py"] + + async def regex_search(self, pattern, file_pattern="**/*", max_results=100): + return [] + + def detect_prompt(self): + return FilesystemAdapter.DETECT_PROMPT + + llm = CountingLlm() + findings, files_read = asyncio.run( + saist_main.generate_findings_with_filesystem_tools_iterations( + scm=FakeScm(), + llm=llm, + filenames=["app.py"], + disable_tools=False, + analysis_skills="", + iterations=5, + max_concurrent=2, + ) + ) + + assert llm.calls == 5 + assert llm.max_active == 2 + assert len(findings) == 5 + assert files_read == set() + + +def test_filesystem_tool_analysis_iterations_run_concurrently(): + class SlowLlm: + async def prompt_structured(self, system_prompt, user_prompt, response_format, tool_fns=None): + await asyncio.sleep(0.05) + return Findings(findings=[]) + + class FakeScm: + async def read_file_contents(self, filename): + return "" + + async def list_files(self): + return ["app.py"] + + async def regex_search(self, pattern, file_pattern="**/*", max_results=100): + return [] + + def detect_prompt(self): + return FilesystemAdapter.DETECT_PROMPT + + started = time.perf_counter() + asyncio.run( + saist_main.generate_findings_with_filesystem_tools_iterations( + scm=FakeScm(), + llm=SlowLlm(), + filenames=["app.py"], + disable_tools=False, + analysis_skills="", + iterations=3, + max_concurrent=3, + ) + ) + elapsed = time.perf_counter() - started + + assert elapsed < 0.12 + + +def test_filesystem_tool_analysis_iterations_cache_each_iteration(tmp_path): + class CountingLlm: + def __init__(self): + self.calls = 0 + + async def prompt_structured(self, system_prompt, user_prompt, response_format, tool_fns=None): + self.calls += 1 + return Findings( + findings=[ + Finding( + file="app.py", + snippet="SECRET", + title=f"Secret {self.calls}", + issue=f"Issue {self.calls}", + recommendation="Fix it.", + validation_steps=["Confirm the issue is reachable."], + cwe="CWE-798", + priority=6, + line_number=self.calls, + ) + ] + ) + + class FakeScm: + async def read_file_contents(self, filename): + return "SECRET = 'dev'\n" + + async def list_files(self): + return ["app.py"] + + async def regex_search(self, pattern, file_pattern="**/*", max_results=100): + return [] + + def detect_prompt(self): + return FilesystemAdapter.DETECT_PROMPT + + llm = CountingLlm() + scm = FakeScm() + cache_dir = tmp_path / "cache" + + findings, _ = asyncio.run( + saist_main.generate_findings_with_filesystem_tools_iterations( + scm=scm, + llm=llm, + filenames=["app.py"], + disable_tools=False, + analysis_skills="", + iterations=2, + max_concurrent=2, + disable_caching=False, + cache_folder=str(cache_dir), + ) + ) + assert llm.calls == 2 + assert [finding.line_number for finding in findings] == [1, 2] + assert sorted(path.name.split("-", 1)[0] for path in cache_dir.iterdir()) == ["1", "2"] + + cached_findings, _ = asyncio.run( + saist_main.generate_findings_with_filesystem_tools_iterations( + scm=scm, + llm=llm, + filenames=["app.py"], + disable_tools=False, + analysis_skills="", + iterations=2, + max_concurrent=2, + disable_caching=False, + cache_folder=str(cache_dir), + ) + ) + assert llm.calls == 2 + assert [finding.line_number for finding in cached_findings] == [1, 2] + + extended_findings, _ = asyncio.run( + saist_main.generate_findings_with_filesystem_tools_iterations( + scm=scm, + llm=llm, + filenames=["app.py"], + disable_tools=False, + analysis_skills="", + iterations=3, + max_concurrent=2, + disable_caching=False, + cache_folder=str(cache_dir), + ) + ) + assert llm.calls == 3 + assert [finding.line_number for finding in extended_findings] == [1, 2, 3] + assert sorted(path.name.split("-", 1)[0] for path in cache_dir.iterdir()) == ["1", "2", "3"] + + +def test_print_coverage_reports_read_percentage(capsys): + saist_main.print_coverage({"app.py"}, ["app.py", "settings.py"]) + + assert "LLM file coverage: 1/2 files read (50.0%)" in capsys.readouterr().out + + +def test_dedupe_findings_keeps_highest_priority_for_same_file_and_line(): + low = Finding.model_validate( + { + "file": "app.py", + "snippet": "danger()", + "title": "Low duplicate", + "issue": "Lower severity issue", + "recommendation": "Fix it.", + "cwe": "CWE-20", + "priority": 4, + "line_number": 10, + } + ) + high = Finding.model_validate( + { + "file": "app.py", + "snippet": "danger()", + "title": "High duplicate", + "issue": "Higher severity issue", + "recommendation": "Fix it now.", + "cwe": "CWE-20", + "priority": 8, + "line_number": 10, + } + ) + different_line = Finding.model_validate( + { + "file": "app.py", + "snippet": "other_danger()", + "title": "Different line", + "issue": "Different issue", + "recommendation": "Fix this too.", + "cwe": "CWE-20", + "priority": 5, + "line_number": 11, + } + ) + + deduped = saist_main.dedupe_findings([low, high, different_line]) + + assert deduped == [high, different_line] + + +def test_get_llm_adapter_applies_thinking_option_to_adapter(): + args = SimpleNamespace( + llm="faike", + llm_model="Fake LLM", + llm_api_key=None, + thinking="xhigh", + ollama_base_uri="http://localhost:11434", + ) + + adapter = asyncio.run(saist_main._get_llm_adapter(args)) + + assert adapter.thinking == "xhigh" + + +def test_get_llm_adapter_builds_azure_foundry_adapter(): + args = SimpleNamespace( + llm="azure-foundry", + llm_model="gpt-5-mini", + llm_api_key="azure-key", + azure_openai_endpoint="https://example.openai.azure.com/openai/v1/", + azure_openai_api_version=None, + thinking="high", + ollama_base_uri="http://localhost:11434", + ) + + adapter = asyncio.run(saist_main._get_llm_adapter(args)) + + assert adapter.model_vendor == "Azure AI Foundry" + assert adapter.model_name == "gpt-5-mini" + assert adapter.thinking == "high" + + +def test_get_llm_adapter_passes_openai_base_uri(): + args = SimpleNamespace( + llm="openai", + llm_model="gpt-5-mini", + llm_api_key="test-key", + openai_base_uri="http://openai.example/v1", + thinking="medium", + ollama_base_uri="http://localhost:11434", + ) + + adapter = asyncio.run(saist_main._get_llm_adapter(args)) + + assert adapter.model_vendor == "OpenAI" + assert str(adapter.model.provider.client.base_url) == "http://openai.example/v1/" + + +def test_dedupe_findings_keeps_first_for_same_priority(): + first = Finding.model_validate( + { + "file": "app.py", + "snippet": "danger()", + "title": "First duplicate", + "issue": "First issue", + "recommendation": "First fix.", + "cwe": "CWE-20", + "priority": 7, + "line_number": 10, + } + ) + second = Finding.model_validate( + { + "file": "app.py", + "snippet": "danger()", + "title": "Second duplicate", + "issue": "Second issue", + "recommendation": "Second fix.", + "cwe": "CWE-20", + "priority": 7, + "line_number": 10, + } + ) + + assert saist_main.dedupe_findings([first, second]) == [first] + + +def test_dedupe_findings_ignores_cwe_for_same_file_and_line(): + first = Finding.model_validate( + { + "file": "app.py", + "snippet": "danger()", + "title": "First duplicate", + "issue": "First issue", + "recommendation": "First fix.", + "cwe": "CWE-20", + "priority": 6, + "line_number": 10, + } + ) + higher_priority_different_cwe = Finding.model_validate( + { + "file": "app.py", + "snippet": "danger()", + "title": "Second duplicate", + "issue": "Second issue", + "recommendation": "Second fix.", + "cwe": "CWE-89", + "priority": 8, + "line_number": 10, + } + ) + + assert saist_main.dedupe_findings([first, higher_priority_different_cwe]) == [ + higher_priority_different_cwe + ] + + +def test_diff_review_comments_are_built_after_dedupe(): + first = Finding.model_validate( + { + "file": "app.py", + "snippet": "danger()", + "title": "First duplicate", + "issue": "First issue", + "recommendation": "First fix.", + "cwe": "CWE-20", + "priority": 5, + "line_number": 10, + } + ) + second = Finding.model_validate( + { + "file": "app.py", + "snippet": "danger()", + "title": "Second duplicate", + "issue": "Second issue", + "recommendation": "Second fix.", + "cwe": "CWE-20", + "priority": 9, + "line_number": 10, + } + ) + + deduped = saist_main.dedupe_findings([first, second]) + comments = saist_main.build_diff_review_comments(deduped, {"app.py": {10: 14}}) + + assert len(comments) == 1 + assert comments[0]["position"] == 13 + assert "Second issue" in comments[0]["body"] + + +def test_filesystem_shallow_scan_generates_missing_skills(tmp_path, monkeypatch): + project_root = tmp_path / "project" + project_root.mkdir() + (project_root / "app.py").write_text("print('hello')\n", encoding="utf-8") + + class FakeLlm: + model_name = "fake-model" + + async def prompt_structured(self, system_prompt, user_prompt, response_format, tool_fns=None): + return response_format.model_validate( + { + "skills": [ + { + "filename": "authorization-model.md", + "content": "# Authorization\nGenerated for test.", + } + ] + } + ) + + class EmptyFilesystemAdapter: + def get_changed_files(self): + return [] + + async def get_file_contents(self, filename): + return "" + + async def list_files(self): + return [] + + async def regex_search(self, pattern, file_pattern="**/*", max_results=100): + return [] + + def create_review(self, comment, review_comments, request_changes): + raise AssertionError("review should not be created when no files are listed") + + args = SimpleNamespace( + SCM="filesystem", + path=str(project_root), + path_for_comparison=None, + llm="faike", + llm_model=None, + thinking="medium", + verbose=0, + generate_skills=False, + skills_path=".saist/skills", + skills_sample_files=5, + skills_sample_bytes=300, + overwrite_skills=False, + disable_skills=False, + skills_max_bytes=60000, + deep=False, + include=None, + exclude=None, + dry_run=False, + disable_tools=False, + skip_line_length_check=False, + max_line_length=1000, + llm_rate_limit=1, + iterations=3, + disable_caching=True, + cache_folder=str(tmp_path / "cache"), + interactive=False, + csv=False, + web=False, + pdf=False, + ci=False, + project_name=None, + ) + + async def fake_get_llm_adapter(parsed_args): + return FakeLlm() + + monkeypatch.setattr(saist_main, "parse_args", lambda: args) + monkeypatch.setattr(saist_main, "_get_llm_adapter", fake_get_llm_adapter) + monkeypatch.setattr(saist_main, "_get_scm_adapter", lambda parsed_args: EmptyFilesystemAdapter()) + + asyncio.run(saist_main.main()) + + skill_file = project_root / ".saist" / "skills" / "authorization-model.md" + assert skill_file.read_text(encoding="utf-8") == "# Authorization\nGenerated for test.\n" + + +def test_process_file_reuses_cache_for_same_skills_and_salts_cache_when_skills_change(tmp_path, monkeypatch): + class CountingLlm: + def __init__(self): + self.calls = 0 + + async def prompt_structured(self, system_prompt, user_prompt, response_format, tool_fns=None): + self.calls += 1 + return Findings(findings=[]) + + class FakeScm: + async def read_file_contents(self, filename): + return "print('same file')\n" + + async def no_sleep(delay): + return None + + monkeypatch.setattr(saist_main.asyncio, "sleep", no_sleep) + + llm = CountingLlm() + scm = FakeScm() + cache_dir = tmp_path / "cache" + cache_dir.mkdir() + + for skills in ("routing skill", "routing skill", "auth skill"): + asyncio.run( + saist_main.process_file( + scm=scm, + llm=llm, + filename="app.py", + patch_text="@@ -0,0 +1 @@\n+print('same file')\n", + disable_tools=True, + disable_caching=False, + cache_folder=str(cache_dir), + analysis_skills=skills, + ) + ) + + cache_files = sorted(path.name for path in cache_dir.iterdir()) + assert llm.calls == 2 + assert len(cache_files) == 2 + assert any(skills_prompt_digest("routing skill") in name for name in cache_files) + assert any(skills_prompt_digest("auth skill") in name for name in cache_files) + + +def test_analyze_single_file_returns_none_when_llm_raises(): + class FailingLlm: + async def prompt_structured(self, system_prompt, user_prompt, response_format, tool_fns=None): + raise RuntimeError("model unavailable") + + class FakeScm: + async def read_file_contents(self, filename): + return "print('hello')\n" + + result = asyncio.run( + saist_main.analyze_single_file( + scm=FakeScm(), + adapter=FailingLlm(), + filename="app.py", + patch_text="@@ -0,0 +1 @@\n+print('hello')\n", + disable_tools=True, + ) + ) + + assert result is None + + +def test_context_from_finding_returns_context_window(): + class FakeScm: + async def read_file_contents(self, filename): + assert filename == "app.py" + return "\n".join(f"line {index}" for index in range(1, 8)) + + finding = Finding.model_validate( + { + "file": "app.py", + "snippet": "line 4", + "title": "Issue", + "issue": "Issue", + "recommendation": "Fix it.", + "cwe": "CWE-20", + "priority": 4, + "line_number": 4, + } + ) + + context, start, end = asyncio.run(saist_main.context_from_finding(FakeScm(), finding, context_size=2)) + + assert start == 2 + assert end == 6 + assert context == "line 2\nline 3\nline 4\nline 5\nline 6" + + +def test_context_from_finding_returns_none_when_file_read_fails(): + class FailingScm: + async def read_file_contents(self, filename): + raise FileNotFoundError(filename) + + finding = Finding.model_validate( + { + "file": "missing.py", + "snippet": "missing", + "title": "Issue", + "issue": "Issue", + "recommendation": "Fix it.", + "cwe": "CWE-20", + "priority": 4, + "line_number": 1, + } + ) + + assert asyncio.run(saist_main.context_from_finding(FailingScm(), finding)) is None + + +def test_generate_summary_from_findings_returns_fallback_when_llm_raises(): + class FailingLlm: + def prompt(self, system_prompt, user_prompt): + raise RuntimeError("summary unavailable") + + finding = Finding.model_validate( + { + "file": "app.py", + "snippet": "danger()", + "title": "Issue", + "issue": "Issue", + "recommendation": "Fix it.", + "cwe": "CWE-20", + "priority": 4, + "line_number": 1, + } + ) + + assert ( + saist_main.generate_summary_from_findings(FailingLlm(), [finding]) + == "Security issues found. Please review the inline comments." + ) + + +def test_generate_summary_uses_scm_specific_prompt(): + class CapturingLlm: + def __init__(self): + self.system_prompt = None + + def prompt(self, system_prompt, user_prompt): + self.system_prompt = system_prompt + return "summary" + + finding = Finding.model_validate( + { + "file": "app.py", + "snippet": "danger()", + "title": "Issue", + "issue": "Issue", + "recommendation": "Fix it.", + "cwe": "CWE-20", + "priority": 4, + "line_number": 1, + } + ) + + llm = CapturingLlm() + assert saist_main.generate_summary_from_findings(llm, [finding], RealGithub.SUMMARY_PROMPT) == "summary" + assert "GitHub pull request diff-based code security review" in llm.system_prompt + + +def test_review_body_includes_validation_steps(): + finding = Finding.model_validate( + { + "file": "app.py", + "snippet": "danger()", + "title": "Issue", + "issue": "Issue", + "recommendation": "Fix it.", + "validation_steps": ["Call the endpoint.", "Confirm data crosses tenant boundaries."], + "cwe": "CWE-20", + "priority": 4, + "line_number": 1, + } + ) + + body = saist_main.build_finding_review_body(finding) + + assert "**Validation Steps:**" in body + assert "1. Call the endpoint." in body + assert "2. Confirm data crosses tenant boundaries." in body diff --git a/tests/test_reportlab_pdf.py b/tests/test_reportlab_pdf.py new file mode 100644 index 0000000..e650a17 --- /dev/null +++ b/tests/test_reportlab_pdf.py @@ -0,0 +1,240 @@ +from models import FindingContext +from reportlab_pdf import ReportLabPdf + + +class FakeLlm: + model_vendor = "Fake AI" + model_name = "fake-model" + + +STATIC_SUMMARY = ( + "SAIST reviewed the example application changes and found three issues across input validation, " + "authorization, and secret handling. The issues below are static sample data used to verify the " + "ReportLab PDF output path." +) + + +def example_findings(): + return [ + FindingContext.model_validate( + { + "file": "app/views.py", + "snippet": "db.execute('select * from users where name = ' + username)", + "title": "SQL injection through string concatenation", + "issue": ( + "User-controlled input is concatenated directly into a SQL query. " + "An attacker could provide crafted input that changes the query structure, reads data for other " + "users, or bypasses application-level checks. This example deliberately uses a longer issue " + "paragraph so the PDF renderer exercises multi-line issue text without letting the panel overlap " + "the section title. The vulnerable construction also makes later review difficult because the " + "query behaviour is hidden inside string assembly rather than expressed through a database API " + "that separates commands from values. If this pattern is copied into nearby handlers, the same " + "weakness can spread across multiple lookup and reporting paths." + ), + "recommendation": ( + "Use parameterized queries for every user-controlled value and keep validation focused on the " + "expected username format before the database call is made. Add a regression test that passes " + "characters commonly used in SQL injection payloads and confirms they are treated as data rather " + "than executable SQL. This longer recommendation verifies that remediation text wraps cleanly in " + "the generated report. Review adjacent database calls for similar string concatenation and move " + "shared query construction into a small helper if the same lookup pattern appears in more than " + "one place. Prefer a fix that is obvious to future maintainers, because defensive query handling " + "is only reliable when it remains easy to spot during code review." + ), + "validation_steps": [ + "Find the lookup_user route or caller and confirm username is controlled by the requester.", + "Submit a username containing a harmless SQL metacharacter and confirm the generated query changes behavior.", + "Verify the same lookup succeeds after replacing concatenation with a parameterized query.", + ], + "cwe": "CWE-89", + "priority": 9, + "line_number": 42, + "context": ( + "def lookup_user(username):\n" + " audit('lookup', username)\n" + " db.execute('select * from users where name = ' + username)\n" + " return db.fetchone()" + ), + "context_start": 40, + "context_end": 43, + } + ), + FindingContext.model_validate( + { + "file": "app/api/admin.py", + "snippet": "return export_customer_records(customer_id)", + "title": "Missing tenant authorization check", + "issue": "The export endpoint reads customer records without checking that the requester owns the tenant.", + "recommendation": "Enforce a tenant ownership check before exporting customer data.", + "validation_steps": [ + "Authenticate as a user from one tenant.", + "Request another tenant's customer export id and confirm records are returned.", + ], + "cwe": "CWE-862", + "priority": 8, + "line_number": 88, + "context": ( + "@route('/admin/customers//export')\n" + "def export_customer(customer_id):\n" + " require_login()\n" + " return export_customer_records(customer_id)" + ), + "context_start": 85, + "context_end": 88, + } + ), + FindingContext.model_validate( + { + "file": "settings.py", + "snippet": "API_TOKEN = 'dev-token-123'", + "title": "Hardcoded API token", + "issue": "A static API token is stored in source code and could be exposed through the repository.", + "recommendation": "Load secrets from a managed secret store or environment variable.", + "validation_steps": ["Inspect the committed settings file and confirm the token value is present."], + "cwe": "CWE-798", + "priority": 5, + "line_number": 12, + "context": ( + "DEBUG = False\n" + "SERVICE_URL = 'https://api.example.test'\n" + "API_TOKEN = 'dev-token-123'\n" + "TIMEOUT_SECONDS = 10" + ), + "context_start": 10, + "context_end": 13, + } + ), + ] + + +def test_reportlab_pdf_writes_pdf_report(tmp_path, monkeypatch): + monkeypatch.chdir(tmp_path) + args = type( + "Args", + (), + { + "pdf_filename": "example-issues.pdf", + "SCM": "filesystem", + "path": "/example/project", + }, + )() + + ReportLabPdf( + llm=FakeLlm(), + project="Example Project", + findings=example_findings(), + comment=STATIC_SUMMARY, + ).run(args) + + pdf_path = tmp_path / "reporting" / "example-issues.pdf" + assert pdf_path.exists() + assert pdf_path.read_bytes().startswith(b"%PDF") + assert pdf_path.stat().st_size > 1000 + + +def test_reportlab_pdf_splits_long_summary_across_pages(tmp_path, monkeypatch): + monkeypatch.chdir(tmp_path) + args = type( + "Args", + (), + { + "pdf_filename": "long-summary.pdf", + "SCM": "filesystem", + "path": "/example/project", + }, + )() + long_summary = "\n\n".join( + f"Executive Summary paragraph {index}. " + "The assessment identified several high-impact findings and includes enough detail to span " + "multiple pages without forcing the text into a single unbreakable table cell." + for index in range(70) + ) + + ReportLabPdf( + llm=FakeLlm(), + project="Example Project", + findings=example_findings(), + comment=long_summary, + ).run(args) + + pdf_path = tmp_path / "reporting" / "long-summary.pdf" + assert pdf_path.exists() + assert pdf_path.read_bytes().startswith(b"%PDF") + + +def test_reportlab_pdf_handles_long_code_lines_and_long_index_values(tmp_path, monkeypatch): + monkeypatch.chdir(tmp_path) + args = type( + "Args", + (), + { + "pdf_filename": "long-code-lines.pdf", + "SCM": "filesystem", + "path": "/example/project", + }, + )() + long_line = "token = '" + ("abcdef1234567890" * 30) + "'" + findings = example_findings() + findings[0].title = "Long code line remains readable in the generated report" + findings[0].file = "app/security/reports/very/deeply/nested/module_with_a_long_filename.py" + findings[0].context = "def load_token():\n" + long_line + "\nreturn token" + findings[0].context_start = 1 + findings[0].context_end = 3 + findings[0].line_number = 2 + + ReportLabPdf( + llm=FakeLlm(), + project="Example Project", + findings=findings, + comment=STATIC_SUMMARY, + ).run(args) + + pdf_path = tmp_path / "reporting" / "long-code-lines.pdf" + assert pdf_path.exists() + assert pdf_path.read_bytes().startswith(b"%PDF") + + +def test_issue_summary_uses_issue_title_file_and_severity_columns(): + report = ReportLabPdf( + llm=FakeLlm(), + project="Example Project", + findings=example_findings(), + comment=STATIC_SUMMARY, + ) + styles = report._styles() + table = report._issue_summary_table(styles) + + headers = [cell.getPlainText() for cell in table._cellvalues[0]] + assert headers == ["Issue ID", "Title", "File", "Severity"] + severity_cells = [row[3].getPlainText() for row in table._cellvalues[1:]] + assert severity_cells == ["Critical", "High", "Medium"] + + +def test_finding_story_includes_validation_steps_section(): + report = ReportLabPdf( + llm=FakeLlm(), + project="Example Project", + findings=example_findings(), + comment=STATIC_SUMMARY, + ) + story = report._finding_story(1, example_findings()[0], report._styles()) + story_text = "\n".join(getattr(item, "getPlainText", lambda: "")() for item in story) + + assert "Validation steps" in story_text + assert "1. Find the lookup_user route" in story_text + + +def test_code_markup_wraps_long_lines_and_context_highlight_stays_light(): + report = ReportLabPdf( + llm=FakeLlm(), + project="Example Project", + findings=example_findings(), + comment=STATIC_SUMMARY, + ) + long_line = "token = '" + ("abcdef1234567890" * 30) + "'" + + assert "
" in report._code_markup(long_line) + + context_table = report._context_table(example_findings()[0], report._styles()) + background_colours = [command[3] for command in context_table._bkgrndcmds if command[0] == "BACKGROUND"] + assert all(str(colour) != "Color(.113725,.227451,.164706,1)" for colour in background_colours) diff --git a/tests/test_scm_tools.py b/tests/test_scm_tools.py new file mode 100644 index 0000000..3488209 --- /dev/null +++ b/tests/test_scm_tools.py @@ -0,0 +1,56 @@ +import asyncio + +from scm import Scm +from scm.adapters import BaseScmAdapter + + +class StubAdapter(BaseScmAdapter): + def create_review(self, comment, review_comments, request_changes): + return None + + def get_changed_files(self): + return [] + + async def get_file_contents(self, file_path: str): + return f"contents for {file_path}" + + +class ToolAdapter(StubAdapter): + async def list_files(self): + return ["app.py", "README.md"] + + async def regex_search(self, pattern: str, file_pattern: str = "**/*", max_results: int = 100): + return [{"filename": "app.py", "line_number": 1, "column": 1, "match": pattern, "line": pattern}] + + +def test_scm_exposes_llm_tool_functions_in_order(): + scm = Scm(ToolAdapter()) + + assert [tool.__name__ for tool in scm.tool_functions()] == [ + "read_file_contents", + "list_files", + "regex_search", + ] + + +def test_scm_delegates_file_tools_to_adapter(): + scm = Scm(ToolAdapter()) + + assert asyncio.run(scm.read_file_contents("app.py")) == "contents for app.py" + assert asyncio.run(scm.list_files()) == ["app.py", "README.md"] + assert asyncio.run(scm.regex_search("needle")) == [ + { + "filename": "app.py", + "line_number": 1, + "column": 1, + "match": "needle", + "line": "needle", + } + ] + + +def test_base_adapter_file_tool_stubs_do_not_error(): + scm = Scm(StubAdapter()) + + assert asyncio.run(scm.list_files()) == [] + assert asyncio.run(scm.regex_search("needle")) == [] diff --git a/tests/test_skills.py b/tests/test_skills.py new file mode 100644 index 0000000..cfa49e7 --- /dev/null +++ b/tests/test_skills.py @@ -0,0 +1,144 @@ +import asyncio + +from util.skills import ( + format_analysis_skills, + generate_skill_files, + load_analysis_skills, + project_root_from_args, + resolve_skills_dir, + skills_prompt_digest, +) + + +def test_load_analysis_skills_reads_markdown_in_order_and_respects_byte_limit(tmp_path): + skills_dir = tmp_path / ".saist" / "skills" + skills_dir.mkdir(parents=True) + (skills_dir / "b.md").write_text("second skill", encoding="utf-8") + (skills_dir / "a.md").write_text("first skill with a long body", encoding="utf-8") + (skills_dir / "empty.md").write_text(" ", encoding="utf-8") + + skills = load_analysis_skills(skills_dir, max_bytes=12) + + assert len(skills) == 1 + assert skills[0].path.name == "a.md" + assert skills[0].content == "first skill" + assert skills[0].truncated is True + + +def test_format_analysis_skills_wraps_guidance_with_source_names(tmp_path): + skills_dir = tmp_path / "skills" + skills_dir.mkdir() + (skills_dir / "routing.md").write_text("# Routing\nUse controllers.", encoding="utf-8") + + prompt_text = format_analysis_skills(load_analysis_skills(skills_dir)) + + assert "Application analysis skills" in prompt_text + assert "routing.md" in prompt_text + assert "Use controllers." in prompt_text + assert "Use this context to validate exploitability and business impact" in prompt_text + assert "not to report generic best-practice advice" in prompt_text + + +def test_project_root_and_skills_dir_resolution(tmp_path): + args = type("Args", (), {"SCM": "filesystem", "path": str(tmp_path)})() + + assert project_root_from_args(args) == tmp_path.resolve() + assert resolve_skills_dir(tmp_path, ".saist/skills") == (tmp_path / ".saist" / "skills").resolve() + assert resolve_skills_dir(tmp_path, str(tmp_path / "absolute")) == tmp_path / "absolute" + + +def test_skills_prompt_digest_changes_with_content(): + assert skills_prompt_digest("auth model") == skills_prompt_digest("auth model") + assert skills_prompt_digest("auth model") != skills_prompt_digest("routing model") + + +def test_generate_skill_files_samples_project_and_sanitizes_filenames(tmp_path): + project_root = tmp_path / "project" + project_root.mkdir() + (project_root / "app.py").write_text("from flask import Flask\napp = Flask(__name__)\n", encoding="utf-8") + (project_root / "package.json").write_text('{"dependencies": {"express": "latest"}}\n', encoding="utf-8") + skills_dir = project_root / ".saist" / "skills" + + class FakeLlm: + def __init__(self): + self.user_prompt = None + + async def prompt_structured(self, system_prompt, user_prompt, response_format, tool_fns=None): + self.user_prompt = user_prompt + return response_format.model_validate( + { + "skills": [ + { + "filename": "../Authorization Model!!", + "content": "# Authorization\nCheck tenant boundaries.", + } + ] + } + ) + + llm = FakeLlm() + result = asyncio.run( + generate_skill_files( + llm=llm, + project_root=project_root, + skills_dir=skills_dir, + max_files=5, + max_file_bytes=300, + ) + ) + + output_file = skills_dir / "authorization-model.md" + assert result.skills_dir == skills_dir + assert result.written == [output_file] + assert result.skipped == [] + assert result.sampled_files == 2 + assert output_file.read_text(encoding="utf-8") == "# Authorization\nCheck tenant boundaries.\n" + assert "authorization-model.md" in llm.user_prompt + assert "package.json" in llm.user_prompt + + +def test_generate_skill_files_preserves_existing_files_unless_overwrite_is_enabled(tmp_path): + project_root = tmp_path / "project" + project_root.mkdir() + (project_root / "routes.py").write_text("routes = []\n", encoding="utf-8") + skills_dir = project_root / ".saist" / "skills" + skills_dir.mkdir(parents=True) + output_file = skills_dir / "application-routing.md" + output_file.write_text("# Existing\nKeep me.\n", encoding="utf-8") + + class FakeLlm: + async def prompt_structured(self, system_prompt, user_prompt, response_format, tool_fns=None): + return response_format.model_validate( + { + "skills": [ + { + "filename": "application-routing.md", + "content": "# New\nReplace me when asked.", + } + ] + } + ) + + skipped = asyncio.run( + generate_skill_files( + llm=FakeLlm(), + project_root=project_root, + skills_dir=skills_dir, + overwrite=False, + ) + ) + assert skipped.written == [] + assert skipped.skipped == [output_file] + assert output_file.read_text(encoding="utf-8") == "# Existing\nKeep me.\n" + + overwritten = asyncio.run( + generate_skill_files( + llm=FakeLlm(), + project_root=project_root, + skills_dir=skills_dir, + overwrite=True, + ) + ) + assert overwritten.written == [output_file] + assert overwritten.skipped == [] + assert output_file.read_text(encoding="utf-8") == "# New\nReplace me when asked.\n"