Drive sync progress + documents list API + prompt heatmap (#59, #60, #94) (#61)
Deploy Beta / unit-tests (push) Successful in 11s
Unit Tests / test (push) Successful in 10s
Deploy Beta / docker (push) Successful in 21s
Deploy Beta / deploy-beta (push) Successful in 48s

## Summary
- **#59** — Persist and expose Drive sync progress (`sync_total`, `sync_processed`, `sync_added`, `sync_updated`, `sync_failed`) during sync so the FE can show a progress bar.
- **#60** — Paginate/search/sort `GET /documents/` with `workspace=personal|company` scope isolation.
- **#94 (API)** — Add `GET /analytics/user_prompt_heatmap/?tz=` weekday × hour bins for user-entered prompts.

## Test plan
- [x] `manage.py test` drive sync / documents / analytics / drive_tasks suites
- [ ] Manual: trigger Drive sync and poll connections → progress fields update while pending
- [ ] Manual: `GET /documents/?workspace=personal&page=1&search=…&ordering=name`
- [ ] Deploy FE companion PR for #92/#93/#94 UI

Closes #59
Closes #60Reviewed-on: #61
This commit was merged in pull request #61.
This commit is contained in:
2026-08-02 04:49:23 -07:00
parent fd12bb972d
commit 57a2350b8e
11 changed files with 430 additions and 15 deletions
+15 -1
View File
@@ -85,8 +85,22 @@ def enqueue_drive_sync(
connection.last_sync_status = DriveConnection.SyncStatus.PENDING connection.last_sync_status = DriveConnection.SyncStatus.PENDING
connection.last_sync_error = "" connection.last_sync_error = ""
connection.sync_total = 0
connection.sync_processed = 0
connection.sync_added = 0
connection.sync_updated = 0
connection.sync_failed = 0
connection.save( connection.save(
update_fields=["last_sync_status", "last_sync_error", "last_modified"] update_fields=[
"last_sync_status",
"last_sync_error",
"sync_total",
"sync_processed",
"sync_added",
"sync_updated",
"sync_failed",
"last_modified",
]
) )
transaction.on_commit( transaction.on_commit(
@@ -0,0 +1,44 @@
# Generated by Django 6.0 on 2026-08-02 11:40
from django.db import migrations, models
class Migration(migrations.Migration):
dependencies = [
("chat_backend", "0029_personal_drive_rag_without_company"),
]
operations = [
migrations.AddField(
model_name="driveconnection",
name="sync_added",
field=models.PositiveIntegerField(default=0),
),
migrations.AddField(
model_name="driveconnection",
name="sync_failed",
field=models.PositiveIntegerField(default=0),
),
migrations.AddField(
model_name="driveconnection",
name="sync_processed",
field=models.PositiveIntegerField(
default=0,
help_text="Files finished in the current/last sync (includes skips).",
),
),
migrations.AddField(
model_name="driveconnection",
name="sync_total",
field=models.PositiveIntegerField(
default=0,
help_text="Remote files discovered for the current/last sync run.",
),
),
migrations.AddField(
model_name="driveconnection",
name="sync_updated",
field=models.PositiveIntegerField(default=0),
),
]
+12
View File
@@ -480,6 +480,18 @@ class DriveConnection(TimeInfoBase):
max_length=16, choices=SyncStatus.choices, default=SyncStatus.NEVER max_length=16, choices=SyncStatus.choices, default=SyncStatus.NEVER
) )
last_sync_error = models.TextField(blank=True, default="") last_sync_error = models.TextField(blank=True, default="")
# Progress for FE progress bar while last_sync_status=pending (#59).
sync_total = models.PositiveIntegerField(
default=0,
help_text="Remote files discovered for the current/last sync run.",
)
sync_processed = models.PositiveIntegerField(
default=0,
help_text="Files finished in the current/last sync (includes skips).",
)
sync_added = models.PositiveIntegerField(default=0)
sync_updated = models.PositiveIntegerField(default=0)
sync_failed = models.PositiveIntegerField(default=0)
is_active = models.BooleanField(default=True) is_active = models.BooleanField(default=True)
class Meta: class Meta:
+5
View File
@@ -298,6 +298,11 @@ class DriveConnectionSerializer(serializers.ModelSerializer):
"last_sync_at", "last_sync_at",
"last_sync_status", "last_sync_status",
"last_sync_error", "last_sync_error",
"sync_total",
"sync_processed",
"sync_added",
"sync_updated",
"sync_failed",
"is_active", "is_active",
"created", "created",
] ]
+55 -2
View File
@@ -346,11 +346,43 @@ def sync_connection(connection: DriveConnection) -> dict[str, Any]:
result: dict[str, Any] = {"added": 0, "updated": 0, "removed": 0, "failed": []} result: dict[str, Any] = {"added": 0, "updated": 0, "removed": 0, "failed": []}
connection.last_sync_status = DriveConnection.SyncStatus.PENDING connection.last_sync_status = DriveConnection.SyncStatus.PENDING
connection.save(update_fields=["last_sync_status", "last_modified"]) connection.sync_total = 0
connection.sync_processed = 0
connection.sync_added = 0
connection.sync_updated = 0
connection.sync_failed = 0
connection.save(
update_fields=[
"last_sync_status",
"sync_total",
"sync_processed",
"sync_added",
"sync_updated",
"sync_failed",
"last_modified",
]
)
def _persist_progress(**extra: Any) -> None:
fields = [
"sync_total",
"sync_processed",
"sync_added",
"sync_updated",
"sync_failed",
"last_modified",
*extra.keys(),
]
for key, value in extra.items():
setattr(connection, key, value)
connection.save(update_fields=list(dict.fromkeys(fields)))
try: try:
access_token = ensure_fresh_token(connection) access_token = ensure_fresh_token(connection)
remote_files = _list_remote_files(connection, access_token) remote_files = _list_remote_files(connection, access_token)
connection.sync_total = len(remote_files)
_persist_progress()
if connection.kind == DriveConnection.Kind.PERSONAL: if connection.kind == DriveConnection.Kind.PERSONAL:
if connection.user_id is None: if connection.user_id is None:
raise DriveSyncError( raise DriveSyncError(
@@ -379,6 +411,8 @@ def sync_connection(connection: DriveConnection) -> dict[str, Any]:
for remote in remote_files: for remote in remote_files:
existing = existing_docs.get(remote.id) existing = existing_docs.get(remote.id)
if existing is not None and existing.remote_etag == remote.etag: if existing is not None and existing.remote_etag == remote.etag:
connection.sync_processed += 1
_persist_progress()
continue continue
try: try:
@@ -387,6 +421,9 @@ def sync_connection(connection: DriveConnection) -> dict[str, Any]:
) )
except DriveSyncError as exc: except DriveSyncError as exc:
result["failed"].append({"file": remote.name, "error": exc.message}) result["failed"].append({"file": remote.name, "error": exc.message})
connection.sync_failed += 1
connection.sync_processed += 1
_persist_progress()
continue continue
if existing is not None: if existing is not None:
@@ -399,6 +436,7 @@ def sync_connection(connection: DriveConnection) -> dict[str, Any]:
existing.save() existing.save()
document = existing document = existing
result["updated"] += 1 result["updated"] += 1
connection.sync_updated += 1
else: else:
document = Document.objects.create( document = Document.objects.create(
workspace=workspace, workspace=workspace,
@@ -411,6 +449,7 @@ def sync_connection(connection: DriveConnection) -> dict[str, Any]:
) )
document.file.save(filename, ContentFile(content), save=True) document.file.save(filename, ContentFile(content), save=True)
result["added"] += 1 result["added"] += 1
connection.sync_added += 1
try: try:
ingest_result = rag_service.add_files_to_store( ingest_result = rag_service.add_files_to_store(
@@ -442,6 +481,10 @@ def sync_connection(connection: DriveConnection) -> dict[str, Any]:
document.sync_error = str(exc) document.sync_error = str(exc)
document.save(update_fields=["sync_error", "last_modified"]) document.save(update_fields=["sync_error", "last_modified"])
result["failed"].append({"file": remote.name, "error": str(exc)}) result["failed"].append({"file": remote.name, "error": str(exc)})
connection.sync_failed += 1
connection.sync_processed += 1
_persist_progress()
for remote_id, document in existing_docs.items(): for remote_id, document in existing_docs.items():
if remote_id not in remote_by_id: if remote_id not in remote_by_id:
@@ -465,7 +508,17 @@ def sync_connection(connection: DriveConnection) -> dict[str, Any]:
finally: finally:
connection.last_sync_at = timezone.now() connection.last_sync_at = timezone.now()
connection.save( connection.save(
update_fields=["last_sync_status", "last_sync_error", "last_sync_at", "last_modified"] update_fields=[
"last_sync_status",
"last_sync_error",
"last_sync_at",
"sync_total",
"sync_processed",
"sync_added",
"sync_updated",
"sync_failed",
"last_modified",
]
) )
return result return result
@@ -50,6 +50,11 @@ class EnqueueDriveSyncTestCase(TransactionTestCase):
connection.refresh_from_db() connection.refresh_from_db()
self.assertEqual(connection.last_sync_status, DriveConnection.SyncStatus.PENDING) self.assertEqual(connection.last_sync_status, DriveConnection.SyncStatus.PENDING)
self.assertEqual(connection.last_sync_error, "") self.assertEqual(connection.last_sync_error, "")
self.assertEqual(connection.sync_total, 0)
self.assertEqual(connection.sync_processed, 0)
self.assertEqual(connection.sync_added, 0)
self.assertEqual(connection.sync_updated, 0)
self.assertEqual(connection.sync_failed, 0)
mock_dispatch.assert_called_once_with(connection_id=self.connection.id) mock_dispatch.assert_called_once_with(connection_id=self.connection.id)
@mock.patch("chat_backend.drive_tasks._dispatch_sync") @mock.patch("chat_backend.drive_tasks._dispatch_sync")
@@ -397,6 +397,11 @@ class SyncConnectionTestCase(TestCase):
self.connection.refresh_from_db() self.connection.refresh_from_db()
self.assertEqual(self.connection.last_sync_status, DriveConnection.SyncStatus.OK) self.assertEqual(self.connection.last_sync_status, DriveConnection.SyncStatus.OK)
self.assertIsNotNone(self.connection.last_sync_at) self.assertIsNotNone(self.connection.last_sync_at)
self.assertEqual(self.connection.sync_total, 1)
self.assertEqual(self.connection.sync_processed, 1)
self.assertEqual(self.connection.sync_added, 1)
self.assertEqual(self.connection.sync_updated, 0)
self.assertEqual(self.connection.sync_failed, 0)
def test_skips_unchanged_file(self): def test_skips_unchanged_file(self):
Document.objects.create( Document.objects.create(
@@ -415,6 +420,9 @@ class SyncConnectionTestCase(TestCase):
self.assertEqual(result["added"], 0) self.assertEqual(result["added"], 0)
self.assertEqual(result["updated"], 0) self.assertEqual(result["updated"], 0)
self.mock_download.assert_not_called() self.mock_download.assert_not_called()
self.connection.refresh_from_db()
self.assertEqual(self.connection.sync_total, 1)
self.assertEqual(self.connection.sync_processed, 1)
def test_updates_file_when_etag_changes(self): def test_updates_file_when_etag_changes(self):
existing = Document.objects.create( existing = Document.objects.create(
@@ -127,3 +127,60 @@ class AdminAnalyticsTestCase(APITestCase):
self.assertEqual(current["range"], [10, 20]) self.assertEqual(current["range"], [10, 20])
self.assertEqual(current["avg"], 15) self.assertEqual(current["avg"], 15)
self.assertEqual(current["median"], 20) self.assertEqual(current["median"], 20)
class UserPromptHeatmapTestCase(APITestCase):
def setUp(self):
self.user = make_user(company=make_company())
self.other = make_user(email="other@example.com", company=make_company("Other"))
self.client.force_authenticate(user=self.user)
self.conversation = make_conversation(user=self.user)
self.url = reverse("analytics_user_prompt_heatmap")
def test_empty_heatmap(self):
response = self.client.get(self.url, {"tz": "UTC"})
self.assertEqual(response.status_code, status.HTTP_200_OK)
self.assertEqual(response.data["total"], 0)
self.assertEqual(response.data["max"], 0)
self.assertEqual(response.data["matrix"], [[0] * 24 for _ in range(7)])
self.assertIsNone(response.data["peak_cell"])
def test_bins_user_prompts_in_utc(self):
# Wednesday 2026-07-15 14:30 UTC → weekday 2, hour 14
created = datetime.datetime(2026, 7, 15, 14, 30, tzinfo=datetime.timezone.utc)
make_prompt(self.conversation, message="a", created=created)
make_prompt(self.conversation, message="b", created=created)
make_prompt(
self.conversation,
message="llm",
user_created=False,
created=created,
)
other_convo = make_conversation(user=self.other)
make_prompt(other_convo, message="other", created=created)
response = self.client.get(self.url, {"tz": "UTC"})
self.assertEqual(response.status_code, status.HTTP_200_OK)
self.assertEqual(response.data["total"], 2)
self.assertEqual(response.data["max"], 2)
self.assertEqual(response.data["matrix"][2][14], 2)
self.assertEqual(response.data["most_active_day"], "Wed")
self.assertEqual(response.data["most_active_hour"], 14)
self.assertEqual(
response.data["peak_cell"],
{"day": "Wed", "hour": 14, "count": 2},
)
def test_applies_timezone(self):
# 2026-07-15 14:30 UTC → America/Chicago CDT (UTC-5) = 09:30 Wed
created = datetime.datetime(2026, 7, 15, 14, 30, tzinfo=datetime.timezone.utc)
make_prompt(self.conversation, message="a", created=created)
response = self.client.get(self.url, {"tz": "America/Chicago"})
self.assertEqual(response.status_code, status.HTTP_200_OK)
self.assertEqual(response.data["tz"], "America/Chicago")
self.assertEqual(response.data["matrix"][2][9], 1)
self.assertEqual(response.data["matrix"][2][14], 0)
@@ -112,9 +112,11 @@ class DocumentUploadViewTestCase(APITestCase):
response = self.client.get(self.url) response = self.client.get(self.url)
self.assertEqual(response.status_code, status.HTTP_200_OK) self.assertEqual(response.status_code, status.HTTP_200_OK)
self.assertEqual(len(response.data), 1) self.assertEqual(response.data["count"], 1)
self.assertIn("test", response.data[0]["file"]) self.assertEqual(response.data["scope"], "company")
self.assertIn("pdf", response.data[0]["file"]) self.assertEqual(len(response.data["results"]), 1)
self.assertIn("test", response.data["results"][0]["file"])
self.assertIn("pdf", response.data["results"][0]["file"])
def test_list_without_workspace_creates_a_default_one(self): def test_list_without_workspace_creates_a_default_one(self):
"""#46: missing workspace is auto-created instead of 404ing the list.""" """#46: missing workspace is auto-created instead of 404ing the list."""
@@ -123,13 +125,50 @@ class DocumentUploadViewTestCase(APITestCase):
response = self.client.get(self.url) response = self.client.get(self.url)
self.assertEqual(response.status_code, status.HTTP_200_OK) self.assertEqual(response.status_code, status.HTTP_200_OK)
self.assertEqual(response.data, []) self.assertEqual(response.data["count"], 0)
self.assertEqual(response.data["results"], [])
self.assertTrue( self.assertTrue(
DocumentWorkspace.objects.filter( DocumentWorkspace.objects.filter(
company=self.company, name="Default" company=self.company, name="Default"
).exists() ).exists()
) )
def test_list_paginates_searches_and_sorts(self):
"""#60: page / search / ordering query params."""
make_document(self.workspace, name="alpha.pdf")
make_document(self.workspace, name="beta.pdf")
make_document(self.workspace, name="gamma.pdf")
page_one = self.client.get(
self.url, {"page": 1, "page_size": 2, "ordering": "name"}
)
self.assertEqual(page_one.status_code, status.HTTP_200_OK)
self.assertEqual(page_one.data["count"], 3)
self.assertEqual(page_one.data["page"], 1)
self.assertEqual(page_one.data["page_size"], 2)
self.assertEqual(len(page_one.data["results"]), 2)
self.assertIn("alpha", page_one.data["results"][0]["file"])
search = self.client.get(self.url, {"search": "beta"})
self.assertEqual(search.status_code, status.HTTP_200_OK)
self.assertEqual(search.data["count"], 1)
self.assertIn("beta", search.data["results"][0]["file"])
def test_list_personal_scope_isolated_from_company(self):
"""#60: company members can list personal workspace separately."""
make_document(self.workspace, name="company.pdf")
personal = make_workspace(user=self.user, name="Personal")
make_document(personal, name="mine.pdf")
company_list = self.client.get(self.url, {"workspace": "company"})
personal_list = self.client.get(self.url, {"workspace": "personal"})
self.assertEqual(company_list.data["count"], 1)
self.assertIn("company", company_list.data["results"][0]["file"])
self.assertEqual(personal_list.data["count"], 1)
self.assertIn("mine", personal_list.data["results"][0]["file"])
self.assertEqual(personal_list.data["scope"], "personal")
class DocumentDetailViewTestCase(APITestCase): class DocumentDetailViewTestCase(APITestCase):
def setUp(self): def setUp(self):
+6
View File
@@ -22,6 +22,7 @@ from .views import (
UserConversationAnalytics, UserConversationAnalytics,
CompanyUsageAnalytics, CompanyUsageAnalytics,
AdminAnalytics, AdminAnalytics,
UserPromptHeatmap,
DocumentWorkspaceView, DocumentWorkspaceView,
DocumentUploadView, DocumentUploadView,
DocumentDetailView, DocumentDetailView,
@@ -103,6 +104,11 @@ urlpatterns = [
name="analytics_company_usage", name="analytics_company_usage",
), ),
path("analytics/admin/", AdminAnalytics.as_view(), name="analytics_admin"), path("analytics/admin/", AdminAnalytics.as_view(), name="analytics_admin"),
path(
"analytics/user_prompt_heatmap/",
UserPromptHeatmap.as_view(),
name="analytics_user_prompt_heatmap",
),
# document urls # document urls
path( path(
"document_workspaces/", "document_workspaces/",
+180 -8
View File
@@ -51,6 +51,7 @@ from chat_backend.services.assistant_identity import ASSISTANT_SYSTEM_PROMPT
from django.utils import timezone from django.utils import timezone
from django.core.files import File from django.core.files import File
from django.core.files.base import ContentFile from django.core.files.base import ContentFile
from django.db.models import Q
import math import math
import datetime import datetime
import pytz import pytz
@@ -68,7 +69,11 @@ from .email_tasks import (
from finance.services.quotas import FeatureNotAllowed, assert_feature_allowed from finance.services.quotas import FeatureNotAllowed, assert_feature_allowed
from .services.llm_service import AsyncLLMService from .services.llm_service import AsyncLLMService
from .services.rag_services import AsyncRAGService from .services.rag_services import AsyncRAGService
from .services.chat_tenant_scope import ensure_workspace_for_user from .services.chat_tenant_scope import (
ensure_company_workspace,
ensure_personal_workspace,
ensure_workspace_for_user,
)
from .services.title_generator import title_generator from .services.title_generator import title_generator
from .services.moderation_classifier import moderation_classifier, ModerationLabel from .services.moderation_classifier import moderation_classifier, ModerationLabel
from .services.prompt_classifier.prompt_classifier import PromptClassifier, PromptType from .services.prompt_classifier.prompt_classifier import PromptClassifier, PromptType
@@ -730,6 +735,76 @@ class CompanyUsageAnalytics(APIView):
return Response(result[::-1], status=status.HTTP_200_OK) return Response(result[::-1], status=status.HTTP_200_OK)
return Response(result[::-1], status=status.HTTP_200_OK)
class UserPromptHeatmap(APIView):
"""Weekday × hour bins of user-entered prompts (local timezone) (#94)."""
DAY_LABELS = ("Mon", "Tue", "Wed", "Thu", "Fri", "Sat", "Sun")
def get(self, request, format="json"):
tz_name = request.query_params.get("tz", "UTC")
try:
user_tz = pytz.timezone(tz_name)
except pytz.UnknownTimeZoneError:
user_tz = pytz.UTC
tz_name = "UTC"
matrix = [[0 for _ in range(24)] for _ in range(7)]
total = 0
prompts = Prompt.objects.filter(
conversation__user=request.user,
conversation__deleted=False,
user_created=True,
).only("created")
for prompt in prompts.iterator():
created = prompt.created
if timezone.is_naive(created):
created = timezone.make_aware(created, datetime.timezone.utc)
local = created.astimezone(user_tz)
matrix[local.weekday()][local.hour] += 1
total += 1
max_count = max((count for row in matrix for count in row), default=0)
most_active_day = None
most_active_hour = None
peak_cell = None
if max_count > 0:
day_totals = [sum(row) for row in matrix]
most_active_day = self.DAY_LABELS[day_totals.index(max(day_totals))]
hour_totals = [sum(matrix[day][hour] for day in range(7)) for hour in range(24)]
most_active_hour = hour_totals.index(max(hour_totals))
peak_day = 0
peak_hour = 0
for day_idx, row in enumerate(matrix):
for hour_idx, count in enumerate(row):
if count > matrix[peak_day][peak_hour]:
peak_day = day_idx
peak_hour = hour_idx
peak_cell = {
"day": self.DAY_LABELS[peak_day],
"hour": peak_hour,
"count": matrix[peak_day][peak_hour],
}
return Response(
{
"tz": tz_name,
"total": total,
"max": max_count,
"days": list(self.DAY_LABELS),
"hours": list(range(24)),
"matrix": matrix,
"most_active_day": most_active_day,
"most_active_hour": most_active_hour,
"peak_cell": peak_cell,
},
status=status.HTTP_200_OK,
)
class AdminAnalytics(APIView): class AdminAnalytics(APIView):
def get(self, request, format="json"): def get(self, request, format="json"):
number_of_months = 3 number_of_months = 3
@@ -835,6 +910,33 @@ class DocumentWorkspaceView(APIView):
class DocumentUploadView(APIView): class DocumentUploadView(APIView):
# permission_classes = [permissions.IsAuthenticated]Z # permission_classes = [permissions.IsAuthenticated]Z
DOCUMENT_SORT_FIELDS = {
"name": "file",
"-name": "-file",
"created": "created",
"-created": "-created",
"date_uploaded": "created",
"-date_uploaded": "-created",
"processed": "processed",
"-processed": "-processed",
"active": "active",
"-active": "-active",
}
def _resolve_list_workspace(self, request):
"""Return (workspace, scope) for personal|company document lists (#60)."""
raw = (request.query_params.get("workspace") or "").strip().lower()
if not raw:
raw = "company" if request.user.company_id else "personal"
if raw == "personal":
return ensure_personal_workspace(request.user), "personal"
if raw == "company":
if not request.user.company_id:
return None, "company"
return ensure_company_workspace(request.user.company), "company"
return False, raw
def get(self, request): def get(self, request):
logger.debug(f"request_3: {request}") logger.debug(f"request_3: {request}")
try: try:
@@ -842,11 +944,54 @@ class DocumentUploadView(APIView):
except FeatureNotAllowed as exc: except FeatureNotAllowed as exc:
return _feature_gate_response(exc) return _feature_gate_response(exc)
workspace = ensure_workspace_for_user(request.user) workspace, scope = self._resolve_list_workspace(request)
serializer = DocumentSerializer( if workspace is False:
Document.objects.filter(workspace=workspace), many=True return Response(
{"error": "workspace must be 'personal' or 'company'."},
status=status.HTTP_400_BAD_REQUEST,
)
if workspace is None:
return Response(
{"error": "Company workspace requires a company membership."},
status=status.HTTP_400_BAD_REQUEST,
)
queryset = Document.objects.filter(workspace=workspace)
search = (request.query_params.get("search") or request.query_params.get("q") or "").strip()
if search:
queryset = queryset.filter(
Q(file__icontains=search) | Q(remote_name__icontains=search)
)
ordering = (request.query_params.get("ordering") or "-created").strip()
order_by = self.DOCUMENT_SORT_FIELDS.get(ordering, "-created")
queryset = queryset.order_by(order_by, "id")
try:
page = max(int(request.query_params.get("page", 1)), 1)
except (TypeError, ValueError):
page = 1
try:
page_size = int(request.query_params.get("page_size", 20))
except (TypeError, ValueError):
page_size = 20
page_size = min(max(page_size, 1), 100)
total = queryset.count()
start = (page - 1) * page_size
end = start + page_size
serializer = DocumentSerializer(queryset[start:end], many=True)
return Response(
{
"count": total,
"page": page,
"page_size": page_size,
"scope": scope,
"results": serializer.data,
},
status=status.HTTP_200_OK,
) )
return Response(serializer.data, status=status.HTTP_200_OK)
def post(self, request): def post(self, request):
logger.debug(f"request: {request}") logger.debug(f"request: {request}")
@@ -856,7 +1001,18 @@ class DocumentUploadView(APIView):
except FeatureNotAllowed as exc: except FeatureNotAllowed as exc:
return _feature_gate_response(exc) return _feature_gate_response(exc)
workspace = ensure_workspace_for_user(request.user) scope = (request.query_params.get("workspace") or request.data.get("workspace") or "").strip().lower()
if scope == "personal":
workspace = ensure_personal_workspace(request.user)
elif scope == "company":
if not request.user.company_id:
return Response(
{"error": "Company workspace requires a company membership."},
status=status.HTTP_400_BAD_REQUEST,
)
workspace = ensure_company_workspace(request.user.company)
else:
workspace = ensure_workspace_for_user(request.user)
logger.info(request.FILES) logger.info(request.FILES)
file = request.FILES.get("file") file = request.FILES.get("file")
@@ -899,8 +1055,24 @@ class DocumentDetailView(APIView):
# permission_classes = [permissions.IsAuthenticated] # permission_classes = [permissions.IsAuthenticated]
def _get_document(self, request, document_id): def _get_document(self, request, document_id):
workspace = ensure_workspace_for_user(request.user) user = request.user
return Document.objects.filter(workspace=workspace, id=document_id).first() if user.company_id:
return (
Document.objects.filter(id=document_id)
.filter(
Q(workspace__user=user, workspace__company__isnull=True)
| Q(
workspace__company_id=user.company_id,
workspace__user__isnull=True,
)
)
.first()
)
return Document.objects.filter(
id=document_id,
workspace__user=user,
workspace__company__isnull=True,
).first()
def get(self, request, document_id): def get(self, request, document_id):
logger.info(f"request: {request}") logger.info(f"request: {request}")