Skip to content

Commit bf34498

Browse files
authored
feat: add progress bars (#67)
1 parent a2fd224 commit bf34498

1 file changed

Lines changed: 79 additions & 36 deletions

File tree

saist/main.py

Lines changed: 79 additions & 36 deletions
Original file line numberDiff line numberDiff line change
@@ -2,39 +2,42 @@
22
import asyncio
33
import logging
44
import os
5-
65
from typing import Optional
76

87
from dotenv import load_dotenv
8+
from latex import Latex
99

1010
from llm.adapters import BaseLlmAdapter
11+
from llm.adapters.anthropic import AnthropicAdapter
12+
from llm.adapters.bedrock import BedrockAdapter
1113
from llm.adapters.deepseek import DeepseekAdapter
14+
from llm.adapters.faike import FaikeAdapter
1215
from llm.adapters.gemini import GeminiAdapter
13-
from llm.adapters.openai import OpenAiAdapter
1416
from llm.adapters.ollama import OllamaAdapter
15-
from llm.adapters.faike import FaikeAdapter
16-
from models import FindingContext, FindingEnriched, Finding, Findings
17-
from llm.adapters.anthropic import AnthropicAdapter
18-
from llm.adapters.bedrock import BedrockAdapter
19-
from web import FindingsServer
17+
from llm.adapters.openai import OpenAiAdapter
18+
19+
from models import Finding, FindingContext, FindingEnriched, Findings
20+
21+
from scm import BaseScmAdapter, Scm
2022
from scm.adapters.filesystem import FilesystemAdapter
21-
from scm import BaseScmAdapter
2223
from scm.adapters.git import GitAdapter
23-
from util.git import parse_unified_diff
24-
from util.filtering import FilterRules
25-
from util.prompts import prompts
2624
from scm.adapters.github import Github
27-
from scm import Scm
25+
2826
from shell import Shell
29-
from latex import Latex
3027

3128
from util.argparsing import parse_args
32-
29+
from util.caching import *
30+
from util.filtering import FilterRules
31+
from util.git import parse_unified_diff
32+
from util.output import print_banner, write_csv
3333
from util.poem import poem
34+
from util.prompts import prompts
3435

35-
from util.output import print_banner, write_csv
36+
from web import FindingsServer
3637

37-
from util.caching import *
38+
from rich.progress import Progress, SpinnerColumn, TextColumn, BarColumn, TaskProgressColumn, TimeRemainingColumn, TimeElapsedColumn, MofNCompleteColumn
39+
from rich.console import Group
40+
from rich.live import Live
3841

3942
prompts = prompts()
4043
load_dotenv(".env")
@@ -232,6 +235,7 @@ async def main():
232235
# 3) Analyze each file in parallel
233236
print("🔍 Analyzing files for security issues...")
234237
max_workers = min(args.llm_rate_limit, len(app_files))
238+
logging.debug(f"{max_workers=}")
235239
all_findings = await generate_findings(scm, llm, app_files, max_workers, args.disable_tools, args.disable_caching, args.cache_folder)
236240

237241
if not all_findings:
@@ -363,22 +367,22 @@ async def main():
363367
if args.ci and len(all_findings) > 0:
364368
exit(1)
365369

366-
async def process_file(scm: Scm, llm, filename, patch_text, semaphore, disable_tools, disable_caching, cache_folder):
367-
async with semaphore:
368-
start = asyncio.get_event_loop().time()
369-
if disable_caching is True:
370+
async def process_file(scm: Scm, llm, filename, patch_text, disable_tools, disable_caching, cache_folder):
371+
start = asyncio.get_event_loop().time()
372+
if disable_caching is True:
373+
result = await analyze_single_file(scm, llm, filename, patch_text, disable_tools)
374+
else:
375+
hash: str = await hash_file(scm, filename)
376+
cache_file = os.path.join(cache_folder, hash + ".json")
377+
if not os.path.exists(cache_file):
370378
result = await analyze_single_file(scm, llm, filename, patch_text, disable_tools)
379+
store_findings_to_cache_file(filename, result, cache_file)
371380
else:
372-
hash: str = await hash_file(scm, filename)
373-
cache_file = os.path.join(cache_folder, hash + ".json")
374-
if not os.path.exists(cache_file):
375-
result = await analyze_single_file(scm, llm, filename, patch_text, disable_tools)
376-
store_findings_to_cache_file(filename, result, cache_file)
377-
else:
378-
result = findings_from_cache_file(cache_file)
379-
elapsed = asyncio.get_event_loop().time() - start
380-
if elapsed < 1:
381-
await asyncio.sleep(1 - elapsed)
381+
result = findings_from_cache_file(cache_file)
382+
elapsed = asyncio.get_event_loop().time() - start
383+
if elapsed < 1:
384+
await asyncio.sleep(1 - elapsed)
385+
382386
return result
383387

384388
async def generate_findings(scm, llm, app_files, max_concurrent, disable_tools, disable_caching, cache_folder):
@@ -388,13 +392,52 @@ async def generate_findings(scm, llm, app_files, max_concurrent, disable_tools,
388392

389393
semaphore = asyncio.Semaphore(max_concurrent)
390394

391-
tasks = [
392-
process_file(scm, llm, filename, patch_text, semaphore, disable_tools, disable_caching, cache_folder)
393-
for filename, patch_text in app_files
394-
]
395+
overall_progress = Progress(
396+
TextColumn("[bold blue]{task.description}"),
397+
BarColumn(),
398+
MofNCompleteColumn(),
399+
TimeElapsedColumn(),
400+
TimeRemainingColumn())
401+
402+
file_progress = Progress(
403+
SpinnerColumn(),
404+
TextColumn("[blue]{task.description}"),
405+
transient=True
406+
)
407+
408+
progress_group = Group(
409+
overall_progress,
410+
file_progress,
411+
)
412+
413+
def task_progress_wrapper(func, overall_progress, overall_task, file_progress, filename, semaphore):
414+
async def sub_func(*args, **kwargs):
415+
async with semaphore:
416+
task_description_text = f"{filename}..."
417+
file_task = file_progress.add_task(description=task_description_text, transient=True)
418+
task_result = await func(*args, **kwargs)
419+
file_progress.remove_task(file_task)
420+
file_progress.refresh()
421+
overall_progress.update(overall_task, advance=1)
422+
return task_result
423+
return sub_func
424+
395425

396-
all_findings = []
397-
results = await asyncio.gather(*tasks)
426+
with Live(progress_group):
427+
tasks = []
428+
overall_task = overall_progress.add_task(f"Analyzing {len(app_files)} files...", total=len(app_files), start=True) # Add a task
429+
for filename, patch_text in app_files:
430+
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)
431+
tasks.append(
432+
wrapper_func
433+
)
434+
435+
all_findings = []
436+
try:
437+
results = await asyncio.gather(*tasks)
438+
finally:
439+
overall_progress.stop()
440+
file_progress.stop()
398441

399442
for result in results:
400443
if result:

0 commit comments

Comments
 (0)