diff --git a/docker-compose.yml b/docker-compose.yml index b493c39..282aa79 100644 --- a/docker-compose.yml +++ b/docker-compose.yml @@ -19,7 +19,7 @@ services: " django: extends: django_env - command: gunicorn stixify.wsgi:application --bind 0.0.0.0:8004 --reload + command: gunicorn stixify.wsgi:application --bind 0.0.0.0:8004 --preload -w 4 ports: - 8004:8004 depends_on: @@ -27,7 +27,7 @@ services: condition: service_started celery: extends: django_env - command: celery -A stixify.worker worker -l INFO + command: celery -A stixify.worker worker -l INFO --autoscale 8,2 --pool=prefork depends_on: - django - redis @@ -39,7 +39,7 @@ services: extends: django_env command: > bash -c " - celery -A stixify.worker beat -l INFO + celery -A stixify.worker.beat beat -l INFO " depends_on: redis: diff --git a/stixify/classifier/tasks.py b/stixify/classifier/tasks.py index 136e580..db0916d 100644 --- a/stixify/classifier/tasks.py +++ b/stixify/classifier/tasks.py @@ -3,7 +3,7 @@ from typing import Any, List import numpy as np -import openai +from .utils import _openai_client import hdbscan import joblib @@ -17,38 +17,6 @@ class ClusteringCancelled(Exception): pass -def _openai_client(): - openai.api_key = os.getenv("OPENAI_API_KEY") - return openai.Client() - - -def compute_embedding_for_document(doc: DocumentEmbedding): - """Fetch a document by id, compute embedding using OpenAI small-3.""" - if not doc.text: - raise ValueError("Document text is empty, cannot compute embedding") - - client = _openai_client() - try: - resp = client.embeddings.create( - input=doc.text, model="text-embedding-3-small", dimensions=512 - ) - vec = resp.data[0].embedding # list of floats - # store as list of floats; `updated_at` is auto-updated by the model - doc.embedding = vec - doc.save(update_fields=["embedding", "updated_at"]) - print(f"Saved embedding for doc {doc.pk}") - except Exception as e: - print(f"Embedding failed for {doc.pk}: {e}") - raise - - -def create_embedding_text(*texts: List[str]) -> str: - """Create a single string to embed from multiple text fields.""" - # simple concat with separator, could be improved with field weighting or truncation - texts = [t.strip() for t in texts if t and t.strip()] - return " | ".join(texts) - - def run_clustering( min_cluster_size: int = settings.CLASSIFIER_MIN_CLUSTER_SIZE, force: bool = False, diff --git a/stixify/classifier/utils.py b/stixify/classifier/utils.py new file mode 100644 index 0000000..b4eaca1 --- /dev/null +++ b/stixify/classifier/utils.py @@ -0,0 +1,43 @@ +import os +from typing import List + +import openai + + +from .models import DocumentEmbedding + + +class ClusteringCancelled(Exception): + pass + + +def _openai_client(): + openai.api_key = os.getenv("OPENAI_API_KEY") + return openai.Client() + + +def compute_embedding_for_document(doc: DocumentEmbedding): + """Fetch a document by id, compute embedding using OpenAI small-3.""" + if not doc.text: + raise ValueError("Document text is empty, cannot compute embedding") + + client = _openai_client() + try: + resp = client.embeddings.create( + input=doc.text, model="text-embedding-3-small", dimensions=512 + ) + vec = resp.data[0].embedding # list of floats + # store as list of floats; `updated_at` is auto-updated by the model + doc.embedding = vec + doc.save(update_fields=["embedding", "updated_at"]) + print(f"Saved embedding for doc {doc.pk}") + except Exception as e: + print(f"Embedding failed for {doc.pk}: {e}") + raise + + +def create_embedding_text(*texts: List[str]) -> str: + """Create a single string to embed from multiple text fields.""" + # simple concat with separator, could be improved with field weighting or truncation + texts = [t.strip() for t in texts if t and t.strip()] + return " | ".join(texts) \ No newline at end of file diff --git a/stixify/web/models.py b/stixify/web/models.py index e06c2b1..d8b4ba6 100644 --- a/stixify/web/models.py +++ b/stixify/web/models.py @@ -10,8 +10,6 @@ from django.core.cache import cache import uuid, typing from stixify.classifier.models import Cluster, DocumentEmbedding -from stixify.classifier.tasks import compute_embedding_for_document, create_embedding_text -import txt2stix, txt2stix.extractions from django.core.exceptions import ValidationError from datetime import UTC, datetime, timezone from django.utils import timezone as dj_timezone @@ -21,9 +19,11 @@ from dogesec_commons.stixifier.models import Profile from dogesec_commons.identity.models import Identity -from sklearn.metrics.pairwise import cosine_similarity from pgvector.django import CosineDistance - +from stixify.classifier.utils import ( + compute_embedding_for_document, + create_embedding_text, +) if typing.TYPE_CHECKING: from .. import settings @@ -35,8 +35,6 @@ def validate_extractor(types, name): pass - - class TLP_Levels(models.TextChoices): RED = "red" AMBER_STRICT = "amber+strict" @@ -102,7 +100,6 @@ class Meta: abstract = True - def upload_to_func(instance: 'File|FileImage', filename): if isinstance(instance, FileImage): instance = instance.report @@ -206,7 +203,7 @@ def process_mode(self): return self.mode def set_txt2stix_data(self, txt2stix_data): - from txt2stix.txt2stix import Txt2StixData + from txt2stix.utils import Txt2StixData if txt2stix_data is None: return @@ -233,6 +230,7 @@ def set_txt2stix_data(self, txt2stix_data): ) def similar_posts(file, visible_to=None): + if not file.embedding: return [] @@ -253,15 +251,11 @@ def similar_posts(file, visible_to=None): continue if len(results) >= 5: break - similarity_score = cosine_similarity( - file.embedding.embedding.reshape(1, -1), - sfile.embedding.embedding.reshape(1, -1), - )[0][0] results.append( { "id": sfile.id, "name": sfile.name, # or get from related file - "score": similarity_score, + "score": 1 - sfile.distance, "tlp_level": sfile.tlp_level, "owner": sfile.identity_id, "added": sfile.created, diff --git a/stixify/web/serializers.py b/stixify/web/serializers.py index bd5116c..b25f64f 100644 --- a/stixify/web/serializers.py +++ b/stixify/web/serializers.py @@ -14,7 +14,6 @@ import file2txt.parsers.core as f2t_core from rest_framework.exceptions import ValidationError from django.utils.translation import gettext_lazy -from stixify.worker import pdf_converter from django.core.files.base import ContentFile import tempfile from pathlib import Path @@ -156,6 +155,8 @@ def create(self, validated_data): # Handle mhtml-pdf conversion if mode == 'mhtml-pdf': + from stixify.worker import pdf_converter + # Save uploaded file to temporary location with tempfile.NamedTemporaryFile(delete=True, suffix='.mhtml') as temp_file: for chunk in uploaded_file.chunks(): diff --git a/stixify/worker/beat.py b/stixify/worker/beat.py new file mode 100644 index 0000000..5505307 --- /dev/null +++ b/stixify/worker/beat.py @@ -0,0 +1,14 @@ +from datetime import timedelta + +from celery import Celery + + +app = Celery("stixify-beat") +app.config_from_object("os:environ", namespace="CELERY") + +app.conf.beat_schedule = { + "auto_refresh_statistics_data": { + "task": "stixify.worker.tasks.auto_refresh_statistics_data", + "schedule": timedelta(minutes=10), + } +} diff --git a/stixify/worker/celery.py b/stixify/worker/celery.py index c0b98f9..193757b 100644 --- a/stixify/worker/celery.py +++ b/stixify/worker/celery.py @@ -1,4 +1,3 @@ -from datetime import timedelta import os from celery import Celery # Set the default Django settings module for the 'celery' program. @@ -10,12 +9,11 @@ app.config_from_object('os:environ', namespace='CELERY') +app.conf.imports = ( + "stixify.worker.process_post", + "stixify.classifier.tasks", + "stixify.worker.tasks", +) + # Load task modules from all registered Django apps. app.autodiscover_tasks() - -app.conf.beat_schedule = { - "auto_refresh_statistics_data": { - "task": "stixify.worker.tasks.auto_refresh_statistics_data", - "schedule": timedelta(minutes=10), - } -} diff --git a/stixify/worker/process_post.py b/stixify/worker/process_post.py new file mode 100644 index 0000000..70d6f4e --- /dev/null +++ b/stixify/worker/process_post.py @@ -0,0 +1,238 @@ +import logging +import time + +from django.conf import settings +from django.core.cache import cache +from django.core.files.base import File as DjangoFile +from django.db import transaction +from dogesec_commons.stixifier.models import Profile +from dogesec_commons.stixifier.stixifier import ReportProperties, StixifyProcessor +from txt2stix.utils import Txt2StixData + +from stixify.web import models +from stixify.web.models import File, Job +from stixify.worker import pdf_converter + + +ARANGO_UPLOAD_COUNTER_KEY = "arango_upload_active_count" +MAX_CONCURRENT_UPLOADS = 1 +LOCK_TIMEOUT = 300 + + +def acquire_upload_lock(job_id, wait_timeout=LOCK_TIMEOUT): + lock_key = f"arango_upload_lock:{job_id}" + logging.info(f"Attempting to acquire upload lock: {lock_key}") + lock_acquired_at = cache.get(lock_key) + + if lock_acquired_at is not None: + raise RuntimeError(f"Upload lock already held for job {job_id}") + + start_time = time.time() + while True: + active_count = cache.get(ARANGO_UPLOAD_COUNTER_KEY, 0) + if active_count < MAX_CONCURRENT_UPLOADS: + cache.set(ARANGO_UPLOAD_COUNTER_KEY, active_count + 1, LOCK_TIMEOUT) + cache.set(lock_key, time.time(), LOCK_TIMEOUT) + logging.info(f"Acquired upload lock for job {job_id} (active: {active_count + 1}/{MAX_CONCURRENT_UPLOADS})") + return + + if time.time() - start_time > wait_timeout: + raise TimeoutError(f"Timeout waiting for arango upload slot after {wait_timeout}s") + + time.sleep(0.1) + + +def release_upload_lock(job_id): + lock_key = f"arango_upload_lock:{job_id}" + if cache.get(lock_key) is not None: + cache.delete(lock_key) + active_count = cache.get(ARANGO_UPLOAD_COUNTER_KEY, 0) + if active_count > 0: + cache.set(ARANGO_UPLOAD_COUNTER_KEY, active_count - 1, LOCK_TIMEOUT) + logging.info(f"Released upload lock for job {job_id} (active: {max(0, active_count - 1)}/{MAX_CONCURRENT_UPLOADS})") + + +def _object_value_backup(file_id): + return list( + models.ObjectValue.objects.filter(file_id=file_id).values( + "id", + "stix_id", + "type", + "knowledgebase", + "values", + "created", + "modified", + "is_dupe", + ) + ) + + +def _restore_object_values(file_id, backup): + models.ObjectValue.objects.filter(file_id=file_id).delete() + models.ObjectValue.objects.bulk_create( + [models.ObjectValue(file_id=file_id, **values) for values in backup] + ) + + +def _update_reprocess_progress(job, file_id, error=None): + progress = job.extra["progress"] + if error is None: + progress["processed_items"] += 1 + else: + progress["failed_processes"] += 1 + progress["errors"].append( + {"file_id": str(file_id), "message": error} + ) + if progress["failed_processes"] >= settings.REPROCESS_MAX_FAILED_PROCESSES: + progress["stopped_early"] = True + progress["stop_reason"] = "failure_limit_reached" + progress["unprocessed_items"] = max( + 0, + progress["total_items"] + - progress["processed_items"] + - progress["failed_processes"], + ) + + +def _process_file(processor, job, file): + skip_extraction = bool((job.extra or {}).get("skip_extraction")) + is_reprocess = job.type == models.JobType.REPROCESS_FILES + + if is_reprocess and skip_extraction: + processor.output_md = file.markdown_file.open().read().decode() + if not file.txt2stix_data: + raise Exception("no existing extraction data to use for reprocess with skip_extraction=true") + txt2stix_data = Txt2StixData.model_validate(file.txt2stix_data) + processor.txt2stix(txt2stix_data) + else: + logging.info(f"running file2txt on {processor.task_name}") + processor.file2txt() + logging.info(f"running txt2stix on {processor.task_name}") + processor.txt2stix() + + processor.write_bundle(processor.bundler) + + +def process_post_impl(job_id, file_id=None, *args): + job = Job.objects.get(id=job_id) + detached_reprocess = ( + job.type == models.JobType.REPROCESS_FILES and file_id is not None + ) + if detached_reprocess: + progress = job.extra["progress"] + if progress["failed_processes"] >= settings.REPROCESS_MAX_FAILED_PROCESSES: + return job_id + progress["current_file_id"] = str(file_id) + progress["current_index"] = ( + progress["processed_items"] + progress["failed_processes"] + ) + file = None + else: + file = job.file + object_values_backup = None + try: + if detached_reprocess: + file = File.objects.get(pk=file_id) + job.state = models.JobState.PROCESSING + job.save() + processing_profile = file.profile + if job.type == models.JobType.REPROCESS_FILES and (job.extra or {}).get( + "profile_id" + ): + processing_profile = Profile.objects.get(pk=job.extra["profile_id"]) + processor = StixifyProcessor( + file.process_file, + processing_profile, + job_id=job.id, + file2txt_mode=file.process_mode, + report_id=file.id, + ) + external_refs = [ + dict( + source_name="stixify_profile_id", + external_id=str(processing_profile.id), + ) + ] + for source in file.sources or []: + source_ref = dict(source_name="stixify_source") + if source.startswith("http://") or source.startswith("https://"): + source_ref.update(url=source) + else: + source_ref.update(description=source) + external_refs.append(source_ref) + + report_props = ReportProperties( + name=file.name, + identity=file.identity.identity, + tlp_level=file.tlp_level, + confidence=file.confidence, + labels=file.labels, + created=file.created, + kwargs=dict( + external_references=external_refs, + admiralty_source_reliability=file.admiralty_source_reliability, + admiralty_information_credibility=file.admiralty_information_credibility, + pap_level=file.pap_level, + ), + ) + processor.setup( + report_prop=report_props, extra=dict(_stixify_file_id=str(file.id)) + ) + + _process_file(processor, job, file) + + if job.type == models.JobType.REPROCESS_FILES: + object_values_backup = _object_value_backup(file.id) + models.ObjectValue.objects.filter(file_id=file.id).delete() + + acquire_upload_lock(job.id) + try: + logging.info(f"uploading {processor.task_name} to arangodb via stix2arango") + processor.upload_to_arango() + finally: + release_upload_lock(job.id) + + with transaction.atomic(): + new_profile_id = (job.extra or {}).get("profile_id") + if new_profile_id: + file.profile_id = new_profile_id + file.save(update_fields=["profile"]) + file.set_txt2stix_data(processor.txt2stix_data) + file.create_embedding(include_non_incident=settings.CREATE_EMBEDDING_INCLUDE_NON_INCIDENT) + + if job.type == models.JobType.IMPORT_FILE: + file.markdown_file.save("markdown.md", processor.md_file.open(), save=True) + models.FileImage.objects.filter(report=file).delete() + + for image in processor.md_images: + models.FileImage.objects.create( + report=file, file=DjangoFile(image, image.name), name=image.name + ) + if processing_profile.generate_pdf: + converted_file_path = processor.tmpdir / "converted_pdf.pdf" + pdf_converter.make_conversion(processor.filename, converted_file_path) + file.pdf_file.save( + converted_file_path.name, open(converted_file_path, mode="rb") + ) + file.save(update_fields=['markdown_file', 'pdf_file']) + except Exception as e: + error = str(e) + job.error = "failed to process file" + if error: + job.error += f": {error}" + if object_values_backup is not None and file is not None: + try: + _restore_object_values(file.id, object_values_backup) + except Exception: + logging.exception( + "failed to restore ObjectValue data for File %s", file.id + ) + if detached_reprocess: + _update_reprocess_progress(job, file_id, job.error) + logging.error(job.error) + logging.exception(e) + else: + if detached_reprocess: + _update_reprocess_progress(job, file_id) + job.save() + return job_id diff --git a/stixify/worker/tasks.py b/stixify/worker/tasks.py index 6d1424c..e5c5e48 100644 --- a/stixify/worker/tasks.py +++ b/stixify/worker/tasks.py @@ -1,71 +1,21 @@ from datetime import UTC, datetime -import logging import os from pathlib import Path import profile -import time import uuid from django.utils import timezone -from txt2stix import txt2stixBundler -from stixify.web.models import Job, File +from stixify.web.models import Job from stixify.web import models from celery import chain, shared_task -from dogesec_commons.stixifier.stixifier import StixifyProcessor, ReportProperties -from dogesec_commons.stixifier.models import Profile from stixify.web.values.statistics import build_data_and_add_to_cache -from django.core.files.uploadedfile import InMemoryUploadedFile -from django.core.files.storage import default_storage -from django.core.files.base import File as DjangoFile -from django.core.files.base import File as DjangoFile -from django.db import transaction -from django.core.cache import cache import stix2 -from stixify.worker import helpers, pdf_converter +from stixify.worker import helpers from django.conf import settings -from txt2stix.txt2stix import Txt2StixData POLL_INTERVAL = 1 -ARANGO_UPLOAD_COUNTER_KEY = "arango_upload_active_count" -MAX_CONCURRENT_UPLOADS = 1 -LOCK_TIMEOUT = 300 - - -def acquire_upload_lock(job_id, wait_timeout=LOCK_TIMEOUT): - lock_key = f"arango_upload_lock:{job_id}" - logging.info(f"Attempting to acquire upload lock: {lock_key}") - lock_acquired_at = cache.get(lock_key) - - if lock_acquired_at is not None: - raise RuntimeError(f"Upload lock already held for job {job_id}") - - start_time = time.time() - while True: - active_count = cache.get(ARANGO_UPLOAD_COUNTER_KEY, 0) - if active_count < MAX_CONCURRENT_UPLOADS: - cache.set(ARANGO_UPLOAD_COUNTER_KEY, active_count + 1, LOCK_TIMEOUT) - cache.set(lock_key, time.time(), LOCK_TIMEOUT) - logging.info(f"Acquired upload lock for job {job_id} (active: {active_count + 1}/{MAX_CONCURRENT_UPLOADS})") - return - - if time.time() - start_time > wait_timeout: - raise TimeoutError(f"Timeout waiting for arango upload slot after {wait_timeout}s") - - time.sleep(0.1) - - -def release_upload_lock(job_id): - lock_key = f"arango_upload_lock:{job_id}" - if cache.get(lock_key) is not None: - cache.delete(lock_key) - active_count = cache.get(ARANGO_UPLOAD_COUNTER_KEY, 0) - if active_count > 0: - cache.set(ARANGO_UPLOAD_COUNTER_KEY, active_count - 1, LOCK_TIMEOUT) - logging.info(f"Released upload lock for job {job_id} (active: {max(0, active_count - 1)}/{MAX_CONCURRENT_UPLOADS})") - - def new_task(job: Job): if job.type == models.JobType.REPROCESS_FILES and (job.extra or {}).get("file_ids"): task = chain( @@ -108,191 +58,11 @@ def create_reprocessing_job(file_ids, options: dict = None): new_task(job) return job -def _process_file(processor, job, file): - skip_extraction = bool((job.extra or {}).get("skip_extraction")) - is_reprocess = job.type == models.JobType.REPROCESS_FILES - - if is_reprocess and skip_extraction: - processor.output_md = file.markdown_file.open().read().decode() - if not file.txt2stix_data: - raise Exception("no existing extraction data to use for reprocess with skip_extraction=true") - txt2stix_data = Txt2StixData.model_validate(file.txt2stix_data) - processor.txt2stix(txt2stix_data) - else: - logging.info(f"running file2txt on {processor.task_name}") - processor.file2txt() - logging.info(f"running txt2stix on {processor.task_name}") - processor.txt2stix() - - processor.write_bundle(processor.bundler) - - -def _object_value_backup(file_id): - return list( - models.ObjectValue.objects.filter(file_id=file_id).values( - "id", - "stix_id", - "type", - "knowledgebase", - "values", - "created", - "modified", - "is_dupe", - ) - ) - - -def _restore_object_values(file_id, backup): - models.ObjectValue.objects.filter(file_id=file_id).delete() - models.ObjectValue.objects.bulk_create( - [models.ObjectValue(file_id=file_id, **values) for values in backup] - ) - - -def _update_reprocess_progress(job, file_id, error=None): - progress = job.extra["progress"] - if error is None: - progress["processed_items"] += 1 - else: - progress["failed_processes"] += 1 - progress["errors"].append( - {"file_id": str(file_id), "message": error} - ) - if progress["failed_processes"] >= settings.REPROCESS_MAX_FAILED_PROCESSES: - progress["stopped_early"] = True - progress["stop_reason"] = "failure_limit_reached" - progress["unprocessed_items"] = max( - 0, - progress["total_items"] - - progress["processed_items"] - - progress["failed_processes"], - ) - - @shared_task def process_post(job_id, file_id=None, *args): - job = Job.objects.get(id=job_id) - detached_reprocess = ( - job.type == models.JobType.REPROCESS_FILES and file_id is not None - ) - if detached_reprocess: - progress = job.extra["progress"] - if progress["failed_processes"] >= settings.REPROCESS_MAX_FAILED_PROCESSES: - return job_id - progress["current_file_id"] = str(file_id) - progress["current_index"] = ( - progress["processed_items"] + progress["failed_processes"] - ) - file = None - else: - file = job.file - object_values_backup = None - try: - if detached_reprocess: - file = File.objects.get(pk=file_id) - job.state = models.JobState.PROCESSING - job.save() - processing_profile = file.profile - if job.type == models.JobType.REPROCESS_FILES and (job.extra or {}).get( - "profile_id" - ): - processing_profile = Profile.objects.get(pk=job.extra["profile_id"]) - processor = StixifyProcessor( - file.process_file, - processing_profile, - job_id=job.id, - file2txt_mode=file.process_mode, - report_id=file.id, - ) - external_refs = [ - dict( - source_name="stixify_profile_id", - external_id=str(processing_profile.id), - ) - ] - for source in file.sources or []: - source_ref = dict(source_name="stixify_source") - if source.startswith("http://") or source.startswith("https://"): - source_ref.update(url=source) - else: - source_ref.update(description=source) - external_refs.append(source_ref) - - report_props = ReportProperties( - name=file.name, - identity=file.identity.identity, - tlp_level=file.tlp_level, - confidence=file.confidence, - labels=file.labels, - created=file.created, - kwargs=dict( - external_references=external_refs, - admiralty_source_reliability=file.admiralty_source_reliability, - admiralty_information_credibility=file.admiralty_information_credibility, - pap_level=file.pap_level, - ), - ) - processor.setup( - report_prop=report_props, extra=dict(_stixify_file_id=str(file.id)) - ) - - _process_file(processor, job, file) - - if job.type == models.JobType.REPROCESS_FILES: - object_values_backup = _object_value_backup(file.id) - models.ObjectValue.objects.filter(file_id=file.id).delete() + from stixify.worker.process_post import process_post_impl - acquire_upload_lock(job.id) - try: - logging.info(f"uploading {processor.task_name} to arangodb via stix2arango") - processor.upload_to_arango() - finally: - release_upload_lock(job.id) - - with transaction.atomic(): - new_profile_id = (job.extra or {}).get("profile_id") - if new_profile_id: - file.profile_id = new_profile_id - file.save(update_fields=["profile"]) - file.set_txt2stix_data(processor.txt2stix_data) - file.create_embedding(include_non_incident=settings.CREATE_EMBEDDING_INCLUDE_NON_INCIDENT) - - if job.type == models.JobType.IMPORT_FILE: - file.markdown_file.save("markdown.md", processor.md_file.open(), save=True) - models.FileImage.objects.filter(report=file).delete() - - for image in processor.md_images: - models.FileImage.objects.create( - report=file, file=DjangoFile(image, image.name), name=image.name - ) - if processing_profile.generate_pdf: - converted_file_path = processor.tmpdir / "converted_pdf.pdf" - pdf_converter.make_conversion(processor.filename, converted_file_path) - file.pdf_file.save( - converted_file_path.name, open(converted_file_path, mode="rb") - ) - file.save(update_fields=['markdown_file', 'pdf_file']) - except Exception as e: - error = str(e) - job.error = "failed to process file" - if error: - job.error += f": {error}" - if object_values_backup is not None and file is not None: - try: - _restore_object_values(file.id, object_values_backup) - except Exception: - logging.exception( - "failed to restore ObjectValue data for File %s", file.id - ) - if detached_reprocess: - _update_reprocess_progress(job, file_id, job.error) - logging.error(job.error) - logging.exception(e) - else: - if detached_reprocess: - _update_reprocess_progress(job, file_id) - job.save() - return job_id + return process_post_impl(job_id, file_id, *args) @shared_task diff --git a/stixify/worker/topics.py b/stixify/worker/topics.py index d12043a..f29c635 100644 --- a/stixify/worker/topics.py +++ b/stixify/worker/topics.py @@ -2,7 +2,6 @@ from concurrent.futures import ThreadPoolExecutor, as_completed import logging, typing from django.conf import settings -from stixify.classifier import tasks as classifier_tasks from celery import shared_task from stixify.web import models from django.utils import timezone @@ -85,9 +84,10 @@ def run_topic_embeddings_job( job.save(update_fields=["extra", "completion_time"]) def run_topic_clusters_job(job_id, force=False): + from stixify.classifier.tasks import run_clustering job = models.Job.objects.get(pk=job_id) try: - classifier_tasks.run_clustering( + run_clustering( force=force, workers=settings.CLASSIFIER_CONCURRENCY, ) diff --git a/tests/src/test_process_post.py b/tests/src/test_process_post.py new file mode 100644 index 0000000..b43031a --- /dev/null +++ b/tests/src/test_process_post.py @@ -0,0 +1,527 @@ +import io +from pathlib import Path +from unittest.mock import MagicMock, patch, call +import uuid +import pytest +from stixify.worker.tasks import job_completed_with_error, new_task, process_post +from stixify.web import models +from dogesec_commons.stixifier.stixifier import StixifyProcessor +from dogesec_commons.stixifier.models import Profile +from django.core.files.base import ContentFile +from django.test import override_settings +from txt2stix.txt2stix import Txt2StixData + +from stixify.worker import tasks + + +@pytest.fixture(autouse=True) +def always_eager(celery_eager): + yield + + +@pytest.mark.django_db +def test_new_task(stixify_job): + with ( + patch("stixify.worker.tasks.process_post.run") as mock_process_post, + patch( + "stixify.worker.tasks.job_completed_with_error.run" + ) as mock_job_completed_with_error, + ): + new_task(stixify_job) + mock_process_post.assert_called_once_with(stixify_job.id) + mock_job_completed_with_error.assert_called_once_with(stixify_job.id) + + +@pytest.mark.django_db +def test_new_task_detached_reprocesses_each_file(stixify_file): + file_ids = [str(stixify_file.id), str(stixify_file.id)] + job = models.Job.objects.create( + type=models.JobType.REPROCESS_FILES, + extra={ + "file_ids": file_ids, + "progress": { + "total_items": 2, + "processed_items": 0, + "failed_processes": 0, + "unprocessed_items": 2, + "current_file_id": None, + "current_index": None, + "stopped_early": False, + "stop_reason": None, + "errors": [], + }, + }, + ) + with ( + patch("stixify.worker.tasks.process_post.run") as mock_process_file, + patch( + "stixify.worker.tasks.job_completed_with_error.run" + ) as mock_completed, + ): + new_task(job) + + assert mock_process_file.call_args_list == [ + call(job.id, file_ids[0]), + call(job.id, file_ids[1]), + ] + mock_completed.assert_called_once_with(job.id) + + +@pytest.mark.django_db +def test_process_post_job__fails(stixify_job): + with ( + patch( + "stixify.worker.process_post.StixifyProcessor", side_effect=ValueError + ) as mock_stixify_processor_cls, + ): + process_post.si(stixify_job.id).delay() + stixify_job.refresh_from_db() + assert stixify_job.error == "failed to process file" + + mock_stixify_processor_cls.side_effect = ValueError("some error") + process_post.si(stixify_job.id).delay() + stixify_job.refresh_from_db() + assert stixify_job.error == "failed to process file: some error" + + +@pytest.fixture +def fake_stixifier_processor(tmpdir): + mocked_processor = MagicMock() + mocked_processor.summary = "Summarized post" + mocked_processor.md_file.open.return_value = io.BytesIO(b"Generated MD File") + mocked_processor.incident = None + mocked_processor.txt2stix_data = Txt2StixData.model_validate(fake_txt2stix_data()) + mocked_processor.md_images = [] + mocked_processor.tmpdir = MagicMock() + mocked_processor.filename = "test.md" + return mocked_processor + + +@pytest.fixture +def stixify_reprocess_job(stixify_job): + stixify_job.type = models.JobType.REPROCESS_FILES + stixify_job.extra = {} + stixify_job.save(update_fields=["type", "extra"]) + stixify_job.file.set_txt2stix_data(fake_txt2stix_data()) + return stixify_job + + +@pytest.mark.django_db +def test_process_post_job(stixify_job, fake_stixifier_processor): + file = stixify_job.file + + with ( + patch("stixify.worker.process_post.StixifyProcessor") as mock_stixify_processor_cls, + patch("stixify.worker.process_post.pdf_converter.make_conversion") as mock_convert_pdf, + patch.object(models.File, "create_embedding") as mock_create_embedding, + ): + mock_stixify_processor_cls.return_value = fake_stixifier_processor + process_post.si(stixify_job.id).delay() + stixify_job.refresh_from_db() + file.refresh_from_db() + mock_convert_pdf.assert_called_once() + mock_stixify_processor_cls.assert_called_once() + mock_stixify_processor_cls.return_value.setup.assert_called_once() + assert mock_stixify_processor_cls.return_value.setup.call_args[1][ + "extra" + ] == dict(_stixify_file_id=str(file.id)) + assert file.txt2stix_data["content_check"]["threat_score"] == 8 + assert file.ai_describes_incident == True + assert file.markdown_file.read() == b"Generated MD File" + process_stream: io.BytesIO = mock_stixify_processor_cls.call_args[0][0] + process_stream.seek(0) + assert process_stream.read() == file.file.read() + mock_stixify_processor_cls.assert_called_once_with( + process_stream, + stixify_job.profile, + job_id=stixify_job.id, + file2txt_mode=file.mode, + report_id=file.id, + ) + mock_create_embedding.assert_called_once_with(include_non_incident=False) + + +@pytest.mark.django_db +def test_process_post_mhtml_pdf_mode(stixify_job, fake_stixifier_processor): + stixify_job.refresh_from_db() + file = stixify_job.file + file.mode = "mhtml-pdf" + file.pdf_file = ContentFile(b"PDF content", name="test.pdf") + file.save() + with ( + patch("stixify.worker.process_post.StixifyProcessor") as mock_stixify_processor_cls, + patch("stixify.worker.process_post.pdf_converter.convert_mhtml_to_pdf") as mock_convert_pdf, + ): + mock_stixify_processor_cls.return_value = fake_stixifier_processor + process_post.si(stixify_job.id).delay() + process_stream: io.BytesIO = mock_stixify_processor_cls.call_args[0][0] + process_stream.seek(0) + mock_stixify_processor_cls.assert_called_once_with( + process_stream, + stixify_job.profile, + job_id=stixify_job.id, + file2txt_mode="pdf", + report_id=file.id, + ) + assert process_stream.read() == b"PDF content" + + +@pytest.mark.django_db +def test_process_post_reprocess_skip_extraction_no_existing_data( + stixify_reprocess_job, fake_stixifier_processor +): + file = stixify_reprocess_job.file + stixify_reprocess_job.extra = {"skip_extraction": True} + stixify_reprocess_job.save(update_fields=["extra"]) + file.markdown_file.save("test.md", io.BytesIO(b"test content")) + file.txt2stix_data = None + file.save(update_fields=["markdown_file", "txt2stix_data"]) + + with patch("stixify.worker.process_post.StixifyProcessor") as mock_stixify_processor_cls: + mock_stixify_processor_cls.return_value = fake_stixifier_processor + new_task(stixify_reprocess_job) + stixify_reprocess_job.refresh_from_db() + assert "no existing extraction data" in stixify_reprocess_job.error + assert stixify_reprocess_job.state == models.JobState.FAILED + assert ( + stixify_reprocess_job.file.markdown_file.read() == b"test content" + ), "File should not be removed if reprocess fails" + + +def fake_txt2stix_data(): + return Txt2StixData.model_validate( + dict( + content_check=dict( + threat_score=8, + describes_incident=True, + explanation="some explanation", + incident_classification=["class1", "class2"], + summary="some summary", + ) + ) + ) + + +@pytest.mark.django_db +def test_process_post_reprocess_skip_extraction_uses_existing_data( + stixify_reprocess_job, fake_stixifier_processor +): + file = stixify_reprocess_job.file + file.markdown_file.save("test.md", io.BytesIO(b"test content")) + file.save(update_fields=["markdown_file", "txt2stix_data"]) + stixify_reprocess_job.extra = {"skip_extraction": True} + stixify_reprocess_job.save(update_fields=["extra"]) + + with ( + patch("stixify.worker.process_post.StixifyProcessor") as mock_stixify_processor_cls, + patch("stixify.worker.process_post.pdf_converter.make_conversion") as mock_convert_pdf, + patch.object(models.File, "create_embedding") as mock_create_embedding, + ): + mock_stixify_processor_cls.return_value = fake_stixifier_processor + new_task(stixify_reprocess_job) + fake_stixifier_processor.file2txt.assert_not_called() + fake_stixifier_processor.txt2stix.assert_called_once() + fake_stixifier_processor.write_bundle.assert_called_once() + fake_stixifier_processor.upload_to_arango.assert_called_once() + mock_convert_pdf.assert_not_called() + mock_create_embedding.assert_called_once() + + + +@pytest.mark.django_db +def test_process_post_reprocess_skip_extraction_acquires_lock( + stixify_reprocess_job, fake_stixifier_processor +): + from django.core.cache import cache + from stixify.worker.process_post import ARANGO_UPLOAD_COUNTER_KEY + + file = stixify_reprocess_job.file + file.markdown_file.save("test.md", io.BytesIO(b"test content")) + file.save(update_fields=["markdown_file", "txt2stix_data"]) + stixify_reprocess_job.extra = {"skip_extraction": True} + stixify_reprocess_job.save(update_fields=["extra"]) + + cache.clear() + + with ( + patch("stixify.worker.process_post.StixifyProcessor") as mock_stixify_processor_cls, + patch.object(models.File, "create_embedding") as mock_create_embedding, + ): + mock_stixify_processor_cls.return_value = fake_stixifier_processor + process_post.si(stixify_reprocess_job.id).delay() + + lock_key = f"arango_upload_lock:{stixify_reprocess_job.id}" + assert cache.get(lock_key) is None, "Lock should be released after upload" + assert cache.get(ARANGO_UPLOAD_COUNTER_KEY, 0) == 0, "Counter should be 0 after upload" + fake_stixifier_processor.upload_to_arango.assert_called_once() + + +@pytest.mark.django_db +def test_process_post_concurrent_uploads_limited( + stixify_reprocess_job, fake_stixifier_processor +): + from django.core.cache import cache + from stixify.worker.process_post import ARANGO_UPLOAD_COUNTER_KEY, MAX_CONCURRENT_UPLOADS + + file = stixify_reprocess_job.file + file.markdown_file.save("test.md", io.BytesIO(b"test content")) + file.save(update_fields=["markdown_file", "txt2stix_data"]) + + cache.clear() + + with ( + patch("stixify.worker.process_post.StixifyProcessor") as mock_stixify_processor_cls, + patch.object(models.File, "create_embedding") as mock_create_embedding, + ): + mock_stixify_processor_cls.return_value = fake_stixifier_processor + + cache.set(ARANGO_UPLOAD_COUNTER_KEY, MAX_CONCURRENT_UPLOADS, 300) + + from stixify.worker.process_post import acquire_upload_lock + with pytest.raises(TimeoutError): + acquire_upload_lock(stixify_reprocess_job.id, wait_timeout=0.1) + cache.clear() + + + +@pytest.mark.django_db +def test_process_post_reprocess_with_profile_switch( + stixify_reprocess_job, fake_stixifier_processor, stixifier_profile +): + new_profile = Profile.objects.create( + name="new-test-profile", + extractions=stixifier_profile.extractions, + extract_text_from_image=stixifier_profile.extract_text_from_image, + defang=stixifier_profile.defang, + relationship_mode=stixifier_profile.relationship_mode, + ai_settings_relationships=stixifier_profile.ai_settings_relationships, + ai_settings_extractions=stixifier_profile.ai_settings_extractions, + ai_content_check_provider=stixifier_profile.ai_content_check_provider, + ai_create_attack_flow=stixifier_profile.ai_create_attack_flow, + ) + stixify_reprocess_job.extra = { + "skip_extraction": False, + "profile_id": str(new_profile.pk), + } + stixify_reprocess_job.save(update_fields=["extra"]) + + with ( + patch("stixify.worker.process_post.StixifyProcessor") as mock_stixify_processor_cls, + patch.object(models.File, "create_embedding") as mock_create_embedding, + ): + mock_stixify_processor_cls.return_value = fake_stixifier_processor + process_post.si(stixify_reprocess_job.id).delay() + stixify_reprocess_job.file.refresh_from_db() + fake_stixifier_processor.file2txt.assert_called_once() + fake_stixifier_processor.txt2stix.assert_called_once() + assert str(stixify_reprocess_job.file.profile_id) == str(new_profile.pk) + assert mock_stixify_processor_cls.call_args.args[1] == new_profile + mock_create_embedding.assert_called_once() + + +@pytest.mark.django_db +def test_process_post_with_incident(stixify_job, fake_stixifier_processor, tmpdir): + fake_stixifier_processor.txt2stix_data.content_check.describes_incident = True + fake_stixifier_processor.tmpdir = Path(tmpdir) + + + with ( + patch("stixify.worker.process_post.StixifyProcessor") as mock_stixify_processor_cls, + patch.object(models.File, "create_embedding") as mock_create_embedding, + patch("stixify.worker.process_post.pdf_converter.make_conversion") as mock_convert_pdf, + + ): + mock_convert_pdf.side_effect = lambda input_path, output_path: output_path.write_bytes(b"PDF content") + mock_stixify_processor_cls.return_value = fake_stixifier_processor + new_task(stixify_job) + mock_create_embedding.assert_called_once_with(include_non_incident=False) + mock_convert_pdf.assert_called_once_with("test.md", fake_stixifier_processor.tmpdir/"converted_pdf.pdf") + file = models.File.objects.get(pk=stixify_job.file_id) + assert file.ai_describes_incident is True + assert file.ai_incident_summary == "some explanation" + assert file.ai_incident_classification == ["class1", "class2"] + + +@pytest.mark.parametrize( + "settings_value", + [ + True, + False, + ], +) +@pytest.mark.django_db +def test_process_post__creates_embedding( + stixify_job, fake_stixifier_processor, settings_value, settings +): + settings.CREATE_EMBEDDING_INCLUDE_NON_INCIDENT = settings_value + with ( + patch("stixify.worker.process_post.StixifyProcessor") as mock_stixify_processor_cls, + patch.object(models.File, "create_embedding") as mock_create_embedding, + patch("stixify.worker.process_post.pdf_converter.make_conversion") as mock_convert_pdf, + ): + mock_stixify_processor_cls.return_value = fake_stixifier_processor + process_post.si(stixify_job.id).delay() + + mock_create_embedding.assert_called_once_with( + include_non_incident=settings_value + ) + + +@pytest.mark.django_db +def test_process_post_full(stixify_job): + with patch("stixify.worker.process_post.pdf_converter.make_conversion") as mock_convert_pdf: + mock_convert_pdf.side_effect = lambda input_path, output_path: output_path.write_bytes(b"%PDF-1.4") + process_post.si(stixify_job.id).delay() + file = models.File.objects.get(pk=stixify_job.file_id) + stixify_job.refresh_from_db() + assert stixify_job.error == None, stixify_job.error + assert tuple(file.archived_pdf.read(4)) == (0x25, 0x50, 0x44, 0x46) + + +@pytest.mark.django_db +def test_job_completed_with_error__failed(stixify_job): + stixify_job.error = "failed" + stixify_job.save() + file_id = stixify_job.file.pk + job_completed_with_error(stixify_job.id) + stixify_job.refresh_from_db() + assert stixify_job.file == None + assert stixify_job.state == models.JobState.FAILED + with pytest.raises(models.File.DoesNotExist): + models.File.objects.get(pk=file_id) + assert stixify_job.completion_time != None + + +@pytest.mark.django_db +def test_job_completed_with_error__success(stixify_job): + file_id = uuid.UUID(stixify_job.file.pk) + job_completed_with_error(stixify_job.id) + stixify_job.refresh_from_db() + assert stixify_job.file.pk == file_id + assert stixify_job.state == models.JobState.COMPLETED + assert stixify_job.completion_time != None + + +def detached_reprocess_job(file_ids, **options): + progress = { + "total_items": len(file_ids), + "processed_items": 0, + "failed_processes": 0, + "unprocessed_items": len(file_ids), + "current_file_id": None, + "current_index": None, + "stopped_early": False, + "stop_reason": None, + "errors": [], + } + return models.Job.objects.create( + type=models.JobType.REPROCESS_FILES, + extra={"file_ids": file_ids, "progress": progress, **options}, + ) + + +@pytest.mark.django_db +def test_detached_reprocess_updates_progress( + stixify_file, fake_stixifier_processor +): + job = detached_reprocess_job([str(stixify_file.id)], skip_extraction=False) + with ( + patch("stixify.worker.process_post.StixifyProcessor") as processor_class, + patch.object(models.File, "create_embedding"), + ): + processor_class.return_value = fake_stixifier_processor + process_post(job.id, stixify_file.id) + + job.refresh_from_db() + assert job.extra["progress"] == { + "total_items": 1, + "processed_items": 1, + "failed_processes": 0, + "unprocessed_items": 0, + "current_file_id": str(stixify_file.id), + "current_index": 0, + "stopped_early": False, + "stop_reason": None, + "errors": [], + } + + +@pytest.mark.django_db +@override_settings(REPROCESS_MAX_FAILED_PROCESSES=10) +def test_detached_reprocess_stops_after_ten_failures(stixify_file): + file_ids = [str(stixify_file.id)] * 11 + job = detached_reprocess_job(file_ids, skip_extraction=False) + + with patch( + "stixify.worker.process_post.StixifyProcessor", side_effect=ValueError("bad file") + ) as processor_class: + for file_id in file_ids: + process_post(job.id, file_id) + + job.refresh_from_db() + progress = job.extra["progress"] + assert processor_class.call_count == 10 + assert progress["processed_items"] == 0 + assert progress["failed_processes"] == 10 + assert progress["unprocessed_items"] == 1 + assert progress["stopped_early"] is True + assert progress["stop_reason"] == "failure_limit_reached" + assert len(progress["errors"]) == 10 + assert set(progress["errors"][0]) == {"file_id", "message"} + + job_completed_with_error(job.id) + job.refresh_from_db() + assert job.state == models.JobState.FAILED + assert job.error == "failed to reprocess 10 file(s)" + assert job.extra["progress"]["current_file_id"] is None + + +@pytest.mark.django_db +def test_detached_reprocess_is_failed_when_one_file_fails( + stixify_file, fake_stixifier_processor +): + file_ids = [str(stixify_file.id), str(stixify_file.id)] + job = detached_reprocess_job(file_ids, skip_extraction=False) + + with ( + patch( + "stixify.worker.process_post.StixifyProcessor", + side_effect=[ValueError("bad file"), fake_stixifier_processor], + ), + patch.object(models.File, "create_embedding"), + ): + for file_id in file_ids: + process_post(job.id, file_id) + job_completed_with_error(job.id) + + job.refresh_from_db() + assert job.state == models.JobState.FAILED + assert job.extra["progress"]["processed_items"] == 1 + assert job.extra["progress"]["failed_processes"] == 1 + assert job.extra["progress"]["unprocessed_items"] == 0 + + +@pytest.mark.django_db +def test_reprocess_restores_object_values_after_failure( + stixify_file, fake_stixifier_processor +): + original = models.ObjectValue.objects.create( + file=stixify_file, + stix_id="indicator--11111111-1111-4111-8111-111111111111", + type="indicator", + values={"name": "original"}, + ) + job = detached_reprocess_job([str(stixify_file.id)], skip_extraction=False) + fake_stixifier_processor.upload_to_arango.side_effect = RuntimeError("upload failed") + + with ( + patch("stixify.worker.process_post.StixifyProcessor") as processor_class, + patch.object(models.File, "create_embedding"), + ): + processor_class.return_value = fake_stixifier_processor + process_post(job.id, stixify_file.id) + + restored = models.ObjectValue.objects.get( + file=stixify_file, stix_id=original.stix_id + ) + assert restored.values == {"name": "original"} diff --git a/tests/src/test_tasks.py b/tests/src/test_tasks.py index 7f226e3..3ef60c9 100644 --- a/tests/src/test_tasks.py +++ b/tests/src/test_tasks.py @@ -1,527 +1,18 @@ -import io -from pathlib import Path -from unittest.mock import MagicMock, patch, call -import uuid -import pytest -from stixify.worker.tasks import job_completed_with_error, new_task, process_post -from stixify.web import models -from dogesec_commons.stixifier.stixifier import StixifyProcessor -from dogesec_commons.stixifier.models import Profile -from django.core.files.base import ContentFile -from django.test import override_settings -from txt2stix.txt2stix import Txt2StixData +import sys +from types import ModuleType +from unittest.mock import Mock -from stixify.worker import tasks +from stixify.worker.tasks import process_post -@pytest.fixture(autouse=True) -def always_eager(celery_eager): - yield - - -@pytest.mark.django_db -def test_new_task(stixify_job): - with ( - patch("stixify.worker.tasks.process_post.run") as mock_process_post, - patch( - "stixify.worker.tasks.job_completed_with_error.run" - ) as mock_job_completed_with_error, - ): - new_task(stixify_job) - mock_process_post.assert_called_once_with(stixify_job.id) - mock_job_completed_with_error.assert_called_once_with(stixify_job.id) - - -@pytest.mark.django_db -def test_new_task_detached_reprocesses_each_file(stixify_file): - file_ids = [str(stixify_file.id), str(stixify_file.id)] - job = models.Job.objects.create( - type=models.JobType.REPROCESS_FILES, - extra={ - "file_ids": file_ids, - "progress": { - "total_items": 2, - "processed_items": 0, - "failed_processes": 0, - "unprocessed_items": 2, - "current_file_id": None, - "current_index": None, - "stopped_early": False, - "stop_reason": None, - "errors": [], - }, - }, - ) - with ( - patch("stixify.worker.tasks.process_post.run") as mock_process_file, - patch( - "stixify.worker.tasks.job_completed_with_error.run" - ) as mock_completed, - ): - new_task(job) - - assert mock_process_file.call_args_list == [ - call(job.id, file_ids[0]), - call(job.id, file_ids[1]), - ] - mock_completed.assert_called_once_with(job.id) - - -@pytest.mark.django_db -def test_process_post_job__fails(stixify_job): - with ( - patch( - "stixify.worker.tasks.StixifyProcessor", side_effect=ValueError - ) as mock_stixify_processor_cls, - ): - process_post.si(stixify_job.id).delay() - stixify_job.refresh_from_db() - assert stixify_job.error == "failed to process file" - - mock_stixify_processor_cls.side_effect = ValueError("some error") - process_post.si(stixify_job.id).delay() - stixify_job.refresh_from_db() - assert stixify_job.error == "failed to process file: some error" - - -@pytest.fixture -def fake_stixifier_processor(tmpdir): - mocked_processor = MagicMock() - mocked_processor.summary = "Summarized post" - mocked_processor.md_file.open.return_value = io.BytesIO(b"Generated MD File") - mocked_processor.incident = None - mocked_processor.txt2stix_data = Txt2StixData.model_validate(fake_txt2stix_data()) - mocked_processor.md_images = [] - mocked_processor.tmpdir = MagicMock() - mocked_processor.filename = "test.md" - return mocked_processor - - -@pytest.fixture -def stixify_reprocess_job(stixify_job): - stixify_job.type = models.JobType.REPROCESS_FILES - stixify_job.extra = {} - stixify_job.save(update_fields=["type", "extra"]) - stixify_job.file.set_txt2stix_data(fake_txt2stix_data()) - return stixify_job - - -@pytest.mark.django_db -def test_process_post_job(stixify_job, fake_stixifier_processor): - file = stixify_job.file - - with ( - patch("stixify.worker.tasks.StixifyProcessor") as mock_stixify_processor_cls, - patch("stixify.worker.pdf_converter.make_conversion") as mock_convert_pdf, - patch.object(models.File, "create_embedding") as mock_create_embedding, - ): - mock_stixify_processor_cls.return_value = fake_stixifier_processor - process_post.si(stixify_job.id).delay() - stixify_job.refresh_from_db() - file.refresh_from_db() - mock_convert_pdf.assert_called_once() - mock_stixify_processor_cls.assert_called_once() - mock_stixify_processor_cls.return_value.setup.assert_called_once() - assert mock_stixify_processor_cls.return_value.setup.call_args[1][ - "extra" - ] == dict(_stixify_file_id=str(file.id)) - assert file.txt2stix_data["content_check"]["threat_score"] == 8 - assert file.ai_describes_incident == True - assert file.markdown_file.read() == b"Generated MD File" - process_stream: io.BytesIO = mock_stixify_processor_cls.call_args[0][0] - process_stream.seek(0) - assert process_stream.read() == file.file.read() - mock_stixify_processor_cls.assert_called_once_with( - process_stream, - stixify_job.profile, - job_id=stixify_job.id, - file2txt_mode=file.mode, - report_id=file.id, - ) - mock_create_embedding.assert_called_once_with(include_non_incident=False) - - -@pytest.mark.django_db -def test_process_post_mhtml_pdf_mode(stixify_job, fake_stixifier_processor): - stixify_job.refresh_from_db() - file = stixify_job.file - file.mode = "mhtml-pdf" - file.pdf_file = ContentFile(b"PDF content", name="test.pdf") - file.save() - with ( - patch("stixify.worker.tasks.StixifyProcessor") as mock_stixify_processor_cls, - patch("stixify.worker.pdf_converter.convert_mhtml_to_pdf") as mock_convert_pdf, - ): - mock_stixify_processor_cls.return_value = fake_stixifier_processor - process_post.si(stixify_job.id).delay() - process_stream: io.BytesIO = mock_stixify_processor_cls.call_args[0][0] - process_stream.seek(0) - mock_stixify_processor_cls.assert_called_once_with( - process_stream, - stixify_job.profile, - job_id=stixify_job.id, - file2txt_mode="pdf", - report_id=file.id, - ) - assert process_stream.read() == b"PDF content" - - -@pytest.mark.django_db -def test_process_post_reprocess_skip_extraction_no_existing_data( - stixify_reprocess_job, fake_stixifier_processor -): - file = stixify_reprocess_job.file - stixify_reprocess_job.extra = {"skip_extraction": True} - stixify_reprocess_job.save(update_fields=["extra"]) - file.markdown_file.save("test.md", io.BytesIO(b"test content")) - file.txt2stix_data = None - file.save(update_fields=["markdown_file", "txt2stix_data"]) - - with patch("stixify.worker.tasks.StixifyProcessor") as mock_stixify_processor_cls: - mock_stixify_processor_cls.return_value = fake_stixifier_processor - new_task(stixify_reprocess_job) - stixify_reprocess_job.refresh_from_db() - assert "no existing extraction data" in stixify_reprocess_job.error - assert stixify_reprocess_job.state == models.JobState.FAILED - assert ( - stixify_reprocess_job.file.markdown_file.read() == b"test content" - ), "File should not be removed if reprocess fails" - - -def fake_txt2stix_data(): - return Txt2StixData.model_validate( - dict( - content_check=dict( - threat_score=8, - describes_incident=True, - explanation="some explanation", - incident_classification=["class1", "class2"], - summary="some summary", - ) - ) +def test_process_post_delegates_to_process_post_impl(monkeypatch): + process_post_impl = Mock(return_value="job-id") + process_post_module = ModuleType("stixify.worker.process_post") + process_post_module.process_post_impl = process_post_impl + monkeypatch.setitem( + sys.modules, "stixify.worker.process_post", process_post_module ) - -@pytest.mark.django_db -def test_process_post_reprocess_skip_extraction_uses_existing_data( - stixify_reprocess_job, fake_stixifier_processor -): - file = stixify_reprocess_job.file - file.markdown_file.save("test.md", io.BytesIO(b"test content")) - file.save(update_fields=["markdown_file", "txt2stix_data"]) - stixify_reprocess_job.extra = {"skip_extraction": True} - stixify_reprocess_job.save(update_fields=["extra"]) - - with ( - patch("stixify.worker.tasks.StixifyProcessor") as mock_stixify_processor_cls, - patch("stixify.worker.pdf_converter.make_conversion") as mock_convert_pdf, - patch.object(models.File, "create_embedding") as mock_create_embedding, - ): - mock_stixify_processor_cls.return_value = fake_stixifier_processor - new_task(stixify_reprocess_job) - fake_stixifier_processor.file2txt.assert_not_called() - fake_stixifier_processor.txt2stix.assert_called_once() - fake_stixifier_processor.write_bundle.assert_called_once() - fake_stixifier_processor.upload_to_arango.assert_called_once() - mock_convert_pdf.assert_not_called() - mock_create_embedding.assert_called_once() - - - -@pytest.mark.django_db -def test_process_post_reprocess_skip_extraction_acquires_lock( - stixify_reprocess_job, fake_stixifier_processor -): - from django.core.cache import cache - from stixify.worker.tasks import ARANGO_UPLOAD_COUNTER_KEY - - file = stixify_reprocess_job.file - file.markdown_file.save("test.md", io.BytesIO(b"test content")) - file.save(update_fields=["markdown_file", "txt2stix_data"]) - stixify_reprocess_job.extra = {"skip_extraction": True} - stixify_reprocess_job.save(update_fields=["extra"]) - - cache.clear() - - with ( - patch("stixify.worker.tasks.StixifyProcessor") as mock_stixify_processor_cls, - patch.object(models.File, "create_embedding") as mock_create_embedding, - ): - mock_stixify_processor_cls.return_value = fake_stixifier_processor - process_post.si(stixify_reprocess_job.id).delay() - - lock_key = f"arango_upload_lock:{stixify_reprocess_job.id}" - assert cache.get(lock_key) is None, "Lock should be released after upload" - assert cache.get(ARANGO_UPLOAD_COUNTER_KEY, 0) == 0, "Counter should be 0 after upload" - fake_stixifier_processor.upload_to_arango.assert_called_once() - - -@pytest.mark.django_db -def test_process_post_concurrent_uploads_limited( - stixify_reprocess_job, fake_stixifier_processor -): - from django.core.cache import cache - from stixify.worker.tasks import ARANGO_UPLOAD_COUNTER_KEY, MAX_CONCURRENT_UPLOADS - - file = stixify_reprocess_job.file - file.markdown_file.save("test.md", io.BytesIO(b"test content")) - file.save(update_fields=["markdown_file", "txt2stix_data"]) - - cache.clear() - - with ( - patch("stixify.worker.tasks.StixifyProcessor") as mock_stixify_processor_cls, - patch.object(models.File, "create_embedding") as mock_create_embedding, - ): - mock_stixify_processor_cls.return_value = fake_stixifier_processor - - cache.set(ARANGO_UPLOAD_COUNTER_KEY, MAX_CONCURRENT_UPLOADS, 300) - - from stixify.worker.tasks import acquire_upload_lock - with pytest.raises(TimeoutError): - acquire_upload_lock(stixify_reprocess_job.id, wait_timeout=0.1) - cache.clear() - - - -@pytest.mark.django_db -def test_process_post_reprocess_with_profile_switch( - stixify_reprocess_job, fake_stixifier_processor, stixifier_profile -): - new_profile = Profile.objects.create( - name="new-test-profile", - extractions=stixifier_profile.extractions, - extract_text_from_image=stixifier_profile.extract_text_from_image, - defang=stixifier_profile.defang, - relationship_mode=stixifier_profile.relationship_mode, - ai_settings_relationships=stixifier_profile.ai_settings_relationships, - ai_settings_extractions=stixifier_profile.ai_settings_extractions, - ai_content_check_provider=stixifier_profile.ai_content_check_provider, - ai_create_attack_flow=stixifier_profile.ai_create_attack_flow, - ) - stixify_reprocess_job.extra = { - "skip_extraction": False, - "profile_id": str(new_profile.pk), - } - stixify_reprocess_job.save(update_fields=["extra"]) - - with ( - patch("stixify.worker.tasks.StixifyProcessor") as mock_stixify_processor_cls, - patch.object(models.File, "create_embedding") as mock_create_embedding, - ): - mock_stixify_processor_cls.return_value = fake_stixifier_processor - process_post.si(stixify_reprocess_job.id).delay() - stixify_reprocess_job.file.refresh_from_db() - fake_stixifier_processor.file2txt.assert_called_once() - fake_stixifier_processor.txt2stix.assert_called_once() - assert str(stixify_reprocess_job.file.profile_id) == str(new_profile.pk) - assert mock_stixify_processor_cls.call_args.args[1] == new_profile - mock_create_embedding.assert_called_once() - - -@pytest.mark.django_db -def test_process_post_with_incident(stixify_job, fake_stixifier_processor, tmpdir): - fake_stixifier_processor.txt2stix_data.content_check.describes_incident = True - fake_stixifier_processor.tmpdir = Path(tmpdir) - - - with ( - patch("stixify.worker.tasks.StixifyProcessor") as mock_stixify_processor_cls, - patch.object(models.File, "create_embedding") as mock_create_embedding, - patch("stixify.worker.pdf_converter.make_conversion") as mock_convert_pdf, - - ): - mock_convert_pdf.side_effect = lambda input_path, output_path: output_path.write_bytes(b"PDF content") - mock_stixify_processor_cls.return_value = fake_stixifier_processor - new_task(stixify_job) - mock_create_embedding.assert_called_once_with(include_non_incident=False) - mock_convert_pdf.assert_called_once_with("test.md", fake_stixifier_processor.tmpdir/"converted_pdf.pdf") - file = models.File.objects.get(pk=stixify_job.file_id) - assert file.ai_describes_incident is True - assert file.ai_incident_summary == "some explanation" - assert file.ai_incident_classification == ["class1", "class2"] - - -@pytest.mark.parametrize( - "settings_value", - [ - True, - False, - ], -) -@pytest.mark.django_db -def test_process_post__creates_embedding( - stixify_job, fake_stixifier_processor, settings_value, settings -): - settings.CREATE_EMBEDDING_INCLUDE_NON_INCIDENT = settings_value - with ( - patch("stixify.worker.tasks.StixifyProcessor") as mock_stixify_processor_cls, - patch.object(models.File, "create_embedding") as mock_create_embedding, - patch("stixify.worker.pdf_converter.make_conversion") as mock_convert_pdf, - ): - mock_stixify_processor_cls.return_value = fake_stixifier_processor - process_post.si(stixify_job.id).delay() - - mock_create_embedding.assert_called_once_with( - include_non_incident=settings_value - ) - - -@pytest.mark.django_db -def test_process_post_full(stixify_job): - with patch("stixify.worker.pdf_converter.make_conversion") as mock_convert_pdf: - mock_convert_pdf.side_effect = lambda input_path, output_path: output_path.write_bytes(b"%PDF-1.4") - process_post.si(stixify_job.id).delay() - file = models.File.objects.get(pk=stixify_job.file_id) - stixify_job.refresh_from_db() - assert stixify_job.error == None, stixify_job.error - assert tuple(file.archived_pdf.read(4)) == (0x25, 0x50, 0x44, 0x46) - - -@pytest.mark.django_db -def test_job_completed_with_error__failed(stixify_job): - stixify_job.error = "failed" - stixify_job.save() - file_id = stixify_job.file.pk - job_completed_with_error(stixify_job.id) - stixify_job.refresh_from_db() - assert stixify_job.file == None - assert stixify_job.state == models.JobState.FAILED - with pytest.raises(models.File.DoesNotExist): - models.File.objects.get(pk=file_id) - assert stixify_job.completion_time != None - - -@pytest.mark.django_db -def test_job_completed_with_error__success(stixify_job): - file_id = uuid.UUID(stixify_job.file.pk) - job_completed_with_error(stixify_job.id) - stixify_job.refresh_from_db() - assert stixify_job.file.pk == file_id - assert stixify_job.state == models.JobState.COMPLETED - assert stixify_job.completion_time != None - - -def detached_reprocess_job(file_ids, **options): - progress = { - "total_items": len(file_ids), - "processed_items": 0, - "failed_processes": 0, - "unprocessed_items": len(file_ids), - "current_file_id": None, - "current_index": None, - "stopped_early": False, - "stop_reason": None, - "errors": [], - } - return models.Job.objects.create( - type=models.JobType.REPROCESS_FILES, - extra={"file_ids": file_ids, "progress": progress, **options}, - ) - - -@pytest.mark.django_db -def test_detached_reprocess_updates_progress( - stixify_file, fake_stixifier_processor -): - job = detached_reprocess_job([str(stixify_file.id)], skip_extraction=False) - with ( - patch("stixify.worker.tasks.StixifyProcessor") as processor_class, - patch.object(models.File, "create_embedding"), - ): - processor_class.return_value = fake_stixifier_processor - process_post(job.id, stixify_file.id) - - job.refresh_from_db() - assert job.extra["progress"] == { - "total_items": 1, - "processed_items": 1, - "failed_processes": 0, - "unprocessed_items": 0, - "current_file_id": str(stixify_file.id), - "current_index": 0, - "stopped_early": False, - "stop_reason": None, - "errors": [], - } - - -@pytest.mark.django_db -@override_settings(REPROCESS_MAX_FAILED_PROCESSES=10) -def test_detached_reprocess_stops_after_ten_failures(stixify_file): - file_ids = [str(stixify_file.id)] * 11 - job = detached_reprocess_job(file_ids, skip_extraction=False) - - with patch( - "stixify.worker.tasks.StixifyProcessor", side_effect=ValueError("bad file") - ) as processor_class: - for file_id in file_ids: - process_post(job.id, file_id) - - job.refresh_from_db() - progress = job.extra["progress"] - assert processor_class.call_count == 10 - assert progress["processed_items"] == 0 - assert progress["failed_processes"] == 10 - assert progress["unprocessed_items"] == 1 - assert progress["stopped_early"] is True - assert progress["stop_reason"] == "failure_limit_reached" - assert len(progress["errors"]) == 10 - assert set(progress["errors"][0]) == {"file_id", "message"} - - job_completed_with_error(job.id) - job.refresh_from_db() - assert job.state == models.JobState.FAILED - assert job.error == "failed to reprocess 10 file(s)" - assert job.extra["progress"]["current_file_id"] is None - - -@pytest.mark.django_db -def test_detached_reprocess_is_failed_when_one_file_fails( - stixify_file, fake_stixifier_processor -): - file_ids = [str(stixify_file.id), str(stixify_file.id)] - job = detached_reprocess_job(file_ids, skip_extraction=False) - - with ( - patch( - "stixify.worker.tasks.StixifyProcessor", - side_effect=[ValueError("bad file"), fake_stixifier_processor], - ), - patch.object(models.File, "create_embedding"), - ): - for file_id in file_ids: - process_post(job.id, file_id) - job_completed_with_error(job.id) - - job.refresh_from_db() - assert job.state == models.JobState.FAILED - assert job.extra["progress"]["processed_items"] == 1 - assert job.extra["progress"]["failed_processes"] == 1 - assert job.extra["progress"]["unprocessed_items"] == 0 - - -@pytest.mark.django_db -def test_reprocess_restores_object_values_after_failure( - stixify_file, fake_stixifier_processor -): - original = models.ObjectValue.objects.create( - file=stixify_file, - stix_id="indicator--11111111-1111-4111-8111-111111111111", - type="indicator", - values={"name": "original"}, - ) - job = detached_reprocess_job([str(stixify_file.id)], skip_extraction=False) - fake_stixifier_processor.upload_to_arango.side_effect = RuntimeError("upload failed") - - with ( - patch("stixify.worker.tasks.StixifyProcessor") as processor_class, - patch.object(models.File, "create_embedding"), - ): - processor_class.return_value = fake_stixifier_processor - process_post(job.id, stixify_file.id) - - restored = models.ObjectValue.objects.get( - file=stixify_file, stix_id=original.stix_id - ) - assert restored.values == {"name": "original"} + assert process_post.run("job-id", "file-id", "extra") == "job-id" + process_post_impl.assert_called_once_with("job-id", "file-id", "extra") + assert process_post.name == "stixify.worker.tasks.process_post" diff --git a/tests/src/test_tasks__embeddings.py b/tests/src/test_tasks__embeddings.py index 35e9395..d1b8573 100644 --- a/tests/src/test_tasks__embeddings.py +++ b/tests/src/test_tasks__embeddings.py @@ -135,7 +135,7 @@ def test_run_topic_clusters_job_success(): state=models.JobState.PROCESSING, ) - with patch("stixify.worker.topics.classifier_tasks.run_clustering") as mock_clustering: + with patch("stixify.classifier.tasks.run_clustering") as mock_clustering: run_topic_clusters_job(job.id, force=True) job.refresh_from_db() @@ -156,7 +156,7 @@ def test_run_topic_clusters_job_clustering_cancelled(): ) with patch( - "stixify.worker.topics.classifier_tasks.run_clustering", + "stixify.classifier.tasks.run_clustering", side_effect=classifier_tasks.ClusteringCancelled, ): run_topic_clusters_job(job.id, force=False) diff --git a/tests/src/test_worker_imports.py b/tests/src/test_worker_imports.py new file mode 100644 index 0000000..416e5a9 --- /dev/null +++ b/tests/src/test_worker_imports.py @@ -0,0 +1,84 @@ +import os +import subprocess +import sys +import textwrap + +from dotenv import load_dotenv + + +def _run_in_clean_process(code): + load_dotenv() + env = os.environ.copy() + env.setdefault("DJANGO_SETTINGS_MODULE", "stixify.settings") + subprocess.run( + [sys.executable, "-c", textwrap.dedent(code)], + check=True, + env=env, + ) + + +def test_django_and_task_imports_do_not_load_worker_dependencies(): + _run_in_clean_process( + """ + import django + import sys + from unittest.mock import patch + + with patch("dogesec_commons.objects.db_view_creator.startup_func"): + django.setup() + + import stixify.worker.tasks + + # assert not any( + # name == "txt2stix" or name.startswith("txt2stix.") + # for name in sys.modules + # ) + assert "stixify.worker.process_post" not in sys.modules + assert "stixify.worker.pdf_converter" not in sys.modules + assert "stixify.classifier.tasks" not in sys.modules + assert "sklearn.metrics.pairwise" not in sys.modules + assert "joblib" not in sys.modules + """ + ) + + +def test_beat_app_does_not_load_django_or_worker_dependencies(): + _run_in_clean_process( + """ + import os + import sys + + os.environ.pop("DJANGO_SETTINGS_MODULE", None) + from stixify.worker.beat import app + from django.apps import apps + + assert not apps.ready + assert "stixify.worker.tasks" not in sys.modules + assert "stixify.worker.process_post" not in sys.modules + assert "stixify.classifier.tasks" not in sys.modules + assert "txt2stix" not in sys.modules + assert "joblib" not in sys.modules + assert app.conf.beat_schedule["auto_refresh_statistics_data"]["task"] == ( + "stixify.worker.tasks.auto_refresh_statistics_data" + ) + """ + ) + + +def test_celery_imports_preload_worker_dependencies(): + _run_in_clean_process( + """ + import sys + from unittest.mock import patch + + with patch("dogesec_commons.objects.db_view_creator.startup_func"): + from stixify.worker.celery import app + app.loader.import_default_modules() + + assert "stixify.worker.process_post" in sys.modules + assert "stixify.worker.pdf_converter" in sys.modules + assert "stixify.classifier.tasks" in sys.modules + assert "txt2stix" in sys.modules + assert "joblib" in sys.modules + """ + ) diff --git a/tests/src/views/test_file_view.py b/tests/src/views/test_file_view.py index 30fd816..fd30056 100644 --- a/tests/src/views/test_file_view.py +++ b/tests/src/views/test_file_view.py @@ -158,7 +158,7 @@ def test_create_mhtml_pdf(client, stixifier_profile, api_schema, identity): ) as mock_job_serializer_cls, patch("stixify.web.views.new_task") as mock_new_task, patch( - "stixify.web.serializers.pdf_converter.convert_mhtml_to_pdf", + "stixify.worker.pdf_converter.convert_mhtml_to_pdf", return_value=b"pdf bytes", ) as mock_convert_mhtml_to_pdf, ): @@ -565,6 +565,38 @@ def test_file_similar_files_visible_to_passed_to_similar_posts( Transport.get_st_response(resp) ) + +@pytest.mark.django_db +def test_similar_posts_uses_annotated_cosine_distance(more_files): + dimensions = 512 + more_files[0].embedding = DocumentEmbedding.objects.create( + id=more_files[0].id, + text="seed", + embedding=[1.0] + [0.0] * (dimensions - 1), + ) + more_files[0].save(update_fields=["embedding"]) + more_files[1].embedding = DocumentEmbedding.objects.create( + id=more_files[1].id, + text="same direction", + embedding=[1.0] + [0.0] * (dimensions - 1), + ) + more_files[1].save(update_fields=["embedding"]) + more_files[2].embedding = DocumentEmbedding.objects.create( + id=more_files[2].id, + text="orthogonal", + embedding=[0.0, 1.0] + [0.0] * (dimensions - 2), + ) + more_files[2].save(update_fields=["embedding"]) + + similar = more_files[0].similar_posts() + + assert [item["id"] for item in similar] == [ + uuid.UUID(more_files[1].id), + uuid.UUID(more_files[2].id), + ] + assert similar[0]["score"] == pytest.approx(1.0) + assert similar[1]["score"] == pytest.approx(0.0) + @pytest.mark.django_db def test_file_pdf_no_pdf(client, stixify_file, api_schema): resp = client.get(