Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
12 changes: 12 additions & 0 deletions posthog/temporal/oauth.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,10 +12,22 @@
ARRAY_APP_CLIENT_ID_US = "HCWoE0aRFMYxIxFNTTwkOORn5LBjOt2GVDzwSw5W"
ARRAY_APP_CLIENT_ID_EU = "AIvijgMS0dxKEmr5z6odvRd8Pkh5vts3nPTzgzU9"
ARRAY_APP_CLIENT_ID_DEV = "DC5uRLVbGI02YQ82grxgnK6Qn12SXWpCqdPb60oZ"
POSTHOG_DESKTOP_MOBILE_CLIENT_ID_US = "a5TY7w9IjFYfes6dkPgZe6envclWw3bm2UD8ZTlm"
POSTHOG_DESKTOP_MOBILE_CLIENT_ID_EU = "1A7vO138Fh5sYmJislicN4F5HnttI6urmFttxPDU"
POSTHOG_AI_APP_CLIENT_ID_US = "N6UgOECSl98ag1xajxPphGApQXYEVvJIwzCXotKu"
POSTHOG_AI_APP_CLIENT_ID_EU = "0Lizwa3mFSlBuEEQ8V8FMJlskUXpDuSmoEdhzxyi"
POSTHOG_AI_APP_CLIENT_ID_DEV = "DD2ZLG6a2YEUtpPANSzSiIBPuUryYmbndLnKKUy1"

POSTHOG_DESKTOP_OAUTH_CLIENT_IDS = frozenset(
{
ARRAY_APP_CLIENT_ID_DEV,
ARRAY_APP_CLIENT_ID_EU,
ARRAY_APP_CLIENT_ID_US,
POSTHOG_DESKTOP_MOBILE_CLIENT_ID_EU,
POSTHOG_DESKTOP_MOBILE_CLIENT_ID_US,
}
)

# Every OAuth application sandbox agent tokens are minted under. Tokens for these apps
# are only ever created server-side (never via the consent flow or personal API keys),
# so a request bearing one provably originates from a sandbox run.
Expand Down
77 changes: 69 additions & 8 deletions products/tasks/backend/facade/api.py
Original file line number Diff line number Diff line change
Expand Up @@ -58,6 +58,7 @@
ChannelFeedMessage,
CodeInvite,
CodeInviteRedemption,
ComputeSource as ComputeSourceModel,
SandboxCustomImage,
SandboxEnvironment,
SandboxSession,
Expand All @@ -83,6 +84,7 @@
# Value types (not ORM models), safe to expose. External callers compare against the
# string-valued ``.status`` / ``.environment`` / ``.origin_product`` fields on the DTOs.
TaskRunStatus = TaskRun.Status
ComputeSource = ComputeSourceModel
TaskRunEnvironment = TaskRun.Environment
TaskOriginProduct = Task.OriginProduct
TaskRuntime = Task.Runtime
Expand Down Expand Up @@ -2932,6 +2934,7 @@ def signal_task_run_user_message(
message_id: str | None = None,
actor_slack_user_id: str | None = None,
steer: bool = False,
compute_source: ComputeSource | None = None,
) -> bool | None:
"""Queue a user_message follow-up signal on the run's workflow.

Expand Down Expand Up @@ -2965,11 +2968,13 @@ def signal_task_run_user_message(
logger.warning("Follow-up signal target workflow gone for task run %s", run.id)
return False
raise
record_task_run_user_activity(run.id, team_id)
record_task_run_user_activity(run.id, team_id, compute_source=compute_source)
return True


def record_task_run_user_activity(run_id: str | UUID, team_id: int) -> None:
def record_task_run_user_activity(
run_id: str | UUID, team_id: int, *, compute_source: ComputeSource | None = None
) -> None:
"""Stamp a user message against the run's open sandbox usage sessions.

Best-effort (the ledger swallows its own failures): records last-activity on
Expand All @@ -2980,7 +2985,7 @@ def record_task_run_user_activity(run_id: str | UUID, team_id: int) -> None:
record_task_run_user_activity as _record_user_activity,
)

_record_user_activity(run_id, team_id)
_record_user_activity(run_id, team_id, compute_source=compute_source)


def get_task_run_sandbox_connection(
Expand Down Expand Up @@ -3235,7 +3240,12 @@ def _github_credential_source_extra_state(pr_authorship_mode, github_user_token:


def bootstrap_task_run(
task_id: str | UUID, team_id: int, user_id: int | None, *, validated_data: dict
task_id: str | UUID,
team_id: int,
user_id: int | None,
*,
validated_data: dict,
compute_source: ComputeSource | None = None,
) -> contracts.TaskRunCreateResult | None:
"""Create a task run (without starting execution) from validated bootstrap data.

Expand Down Expand Up @@ -3369,7 +3379,13 @@ def bootstrap_task_run(
logger.info(
"Creating task run for task %s with mode=%s, branch=%s, environment=%s", task.id, mode, branch, environment
)
run = task.create_run(environment=environment, mode=mode, branch=branch, extra_state=extra_state)
run = task.create_run(
environment=environment,
mode=mode,
branch=branch,
extra_state=extra_state,
compute_source=compute_source,
)

if imported_mcp_servers or relayed_mcp_servers:
update_fields = ["updated_at"]
Expand Down Expand Up @@ -3463,7 +3479,13 @@ def check_task_run_startable(run_id: str | UUID, task_id: str | UUID, team_id: i


def start_task_run(
run_id: str | UUID, task_id: str | UUID, team_id: int, user_id: int | None, *, validated_data: dict
run_id: str | UUID,
task_id: str | UUID,
team_id: int,
user_id: int | None,
*,
validated_data: dict,
compute_source: ComputeSource | None = None,
) -> tuple[str, UUID | None]:
"""Apply run-scoped attachments and trigger the cloud workflow for a startable run.

Expand All @@ -3480,6 +3502,8 @@ def start_task_run(
if run is None:
return "not_found", None
task = run.task
run.compute_source = compute_source
run.save(update_fields=["compute_source", "updated_at"])

pending_user_message = validated_data.get("pending_user_message")
pending_user_artifact_ids = validated_data.get("pending_user_artifact_ids") or []
Expand Down Expand Up @@ -3525,7 +3549,12 @@ def start_task_run(


def resume_task_run_in_cloud(
run_id: str | UUID, task_id: str | UUID, team_id: int, user_id: int | None
run_id: str | UUID,
task_id: str | UUID,
team_id: int,
user_id: int | None,
*,
compute_source: ComputeSource | None = None,
) -> tuple[str, contracts.TaskRunDetailDTO | None, str | None]:
"""Resume a run in a cloud sandbox, terminating any prior workflow.

Expand Down Expand Up @@ -3595,6 +3624,8 @@ def resume_task_run_in_cloud(
prior_environment = run.environment
prior_completed_at = run.completed_at
prior_state = dict(run.state or {})
prior_compute_source = run.compute_source
run.compute_source = compute_source
run.prepare_for_cloud_handoff()

logger.info("Resuming task run in cloud", extra={"task_run_id": str(run.id), "task_id": str(run.task_id)})
Expand All @@ -3614,8 +3645,19 @@ def resume_task_run_in_cloud(
run.environment = prior_environment
run.completed_at = prior_completed_at
run.state = prior_state
run.compute_source = prior_compute_source
run.error_message = "Failed to start cloud workflow"
run.save(update_fields=["status", "environment", "completed_at", "state", "error_message", "updated_at"])
run.save(
update_fields=[
"status",
"environment",
"completed_at",
"state",
"compute_source",
"error_message",
"updated_at",
]
)
run.publish_stream_state_event()
return "workflow_failed", None, None

Expand Down Expand Up @@ -4097,6 +4139,25 @@ def create_task(team_id: int, user_id: int | None, *, validated_data: dict) -> c
return _task_detail_to_dto(_task_detail_queryset().get(pk=task.pk))


def create_signal_report_task(
team_id: int, user_id: int | None, *, validated_data: dict
) -> tuple[contracts.TaskDetailDTO, bool]:
from products.signals.backend.models import SignalReport, SignalReportTask # noqa: PLC0415
from products.signals.backend.task_run_artefacts import TASK_RUN_TYPE_IMPLEMENTATION # noqa: PLC0415

report = validated_data["signal_report"]
with transaction.atomic():
SignalReport.objects.select_for_update().get(id=report.id, team_id=team_id)
existing = SignalReportTask.objects.filter(
team_id=team_id,
report_id=report.id,
relationship=TASK_RUN_TYPE_IMPLEMENTATION,
).first()
if existing is not None:
return _task_detail_to_dto(_task_detail_queryset().get(pk=existing.task_id)), False
return create_task(team_id, user_id, validated_data=validated_data), True


def set_task_title(task_id: str | UUID, team_id: int, title: str) -> bool:
"""Set a task's title, team-scoped. For automated relabels — e.g. backfilling a Signals research
task with ``"Research: <report title>"`` once research produces the title. Leaves
Expand Down
36 changes: 27 additions & 9 deletions products/tasks/backend/logic/services/sandbox_usage.py
Original file line number Diff line number Diff line change
Expand Up @@ -25,7 +25,7 @@
import structlog

from products.tasks.backend.logic.services.sandbox import SandboxConfig
from products.tasks.backend.models import SandboxSession, TaskRun
from products.tasks.backend.models import ComputeSource, SandboxSession, TaskRun

logger = structlog.get_logger(__name__)

Expand Down Expand Up @@ -58,16 +58,26 @@ def open_sandbox_session(
with transaction.atomic():
run = (
TaskRun.objects.select_for_update(of=("self",))
.select_related("task")
.only("id", "team_id", "state", "task__origin_product")
.select_related("task", "task__loop")
.only(
"id",
"team_id",
"state",
"compute_source",
"task__origin_product",
"task__loop__internal",
)
.get(id=run_id)
)
state = run.state or {}
loop = run.task.loop
created_at = sandbox_created_at or timezone.now()
shape = {
"team_id": run.team_id,
"task_run_id": run.id,
"origin_product": run.task.origin_product,
"compute_source": run.compute_source,
"loop_internal": loop.internal if loop is not None else None,
"prewarmed": bool(state.get("prewarmed")),
"vm_runtime": config.is_vm,
"cpu_cores": config.cpu_cores,
Expand Down Expand Up @@ -110,19 +120,27 @@ def close_sandbox_session(sandbox_id: str, *, reason: str) -> None:


@_best_effort
def record_task_run_user_activity(run_id: str | UUID, team_id: int) -> None:
def record_task_run_user_activity(
run_id: str | UUID, team_id: int, *, compute_source: ComputeSource | None = None
) -> None:
"""Stamp a user message against the run's open sandbox sessions.

Sets ``last_user_activity_at`` on every message and ``user_attributed_at``
set-if-NULL, so the first message both claims a warm sandbox and self-heals the
race where a claim lands mid-provision (before ``open_sandbox_session`` read the
run state).
Sets ``last_user_activity_at`` on every message. The first message atomically
claims a warm sandbox and fixes its compute source, including when the claim
lands mid-provision.
"""
now = timezone.now()
run_uuid = run_id if isinstance(run_id, UUID) else UUID(run_id)
open_sessions = SandboxSession.objects.for_team(team_id).filter(task_run_id=run_uuid, ended_at__isnull=True)
claim_updates: dict[str, object] = {"user_attributed_at": now}
if compute_source is not None:
claim_updates["compute_source"] = compute_source
claimed = open_sessions.filter(user_attributed_at__isnull=True).update(**claim_updates)
if claimed and compute_source is not None:
TaskRun.objects.filter(id=run_uuid, team_id=team_id, compute_source__isnull=True).update(
compute_source=compute_source
)
open_sessions.update(last_user_activity_at=now)
open_sessions.filter(user_attributed_at__isnull=True).update(user_attributed_at=now)


@dataclass(frozen=True)
Expand Down
63 changes: 58 additions & 5 deletions products/tasks/backend/logic/services/tests/test_sandbox_usage.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,7 +14,7 @@
open_sandbox_session,
record_task_run_user_activity,
)
from products.tasks.backend.models import SandboxSession, Task, TaskRun
from products.tasks.backend.models import ComputeSource, Loop, SandboxSession, Task, TaskRun


def _config(**overrides) -> SandboxConfig:
Expand All @@ -24,23 +24,32 @@ def _config(**overrides) -> SandboxConfig:


class SandboxUsageBase(APIBaseTest):
def _run(self, *, state: dict | None = None) -> TaskRun:
def _run(self, *, state: dict | None = None, compute_source: ComputeSource | None = None) -> TaskRun:
task = Task.objects.create(
team=self.team, title="t", description="", origin_product=Task.OriginProduct.USER_CREATED
team=self.team,
title="t",
description="",
origin_product=Task.OriginProduct.USER_CREATED,
)
return TaskRun.objects.create(
task=task,
team=self.team,
state=state or {},
compute_source=compute_source,
)
return TaskRun.objects.create(task=task, team=self.team, state=state or {})


class TestSandboxSessionWrites(SandboxUsageBase):
def test_open_attributes_cold_runs_immediately(self):
run = self._run()
run = self._run(compute_source=ComputeSource.POSTHOG_DESKTOP)

open_sandbox_session(run_id=run.id, sandbox_id="sb-cold", config=_config())

session = SandboxSession.objects.unscoped().get(sandbox_id="sb-cold")
assert session.team_id == self.team.id
assert session.task_run_id == run.id
assert session.origin_product == Task.OriginProduct.USER_CREATED
assert session.compute_source == ComputeSource.POSTHOG_DESKTOP
assert session.user_attributed_at is not None
assert session.prewarmed is False
assert session.vm_runtime is False
Expand Down Expand Up @@ -90,6 +99,27 @@ def test_open_records_vm_runtime(self):

assert SandboxSession.objects.unscoped().get(sandbox_id="sb-vm").vm_runtime is True

def test_open_snapshots_loop_internal_classification(self):
loop = Loop.objects.unscoped().create(
team=self.team,
name="Internal loop",
instructions="Run",
runtime_adapter="claude",
internal=True,
)
task = Task.objects.create(
team=self.team,
title="Loop task",
description="",
origin_product=Task.OriginProduct.LOOP,
loop=loop,
)
run = TaskRun.objects.create(task=task, team=self.team)

open_sandbox_session(run_id=run.id, sandbox_id="sb-loop", config=_config())

assert SandboxSession.objects.unscoped().get(sandbox_id="sb-loop").loop_internal is True

def test_open_retry_never_regresses_attribution(self):
run = self._run(state={"await_user_message": True})
open_sandbox_session(run_id=run.id, sandbox_id="sb-retry", config=_config())
Expand Down Expand Up @@ -174,6 +204,29 @@ def test_facade_signal_attributes_claimed_warm_run(self):

assert SandboxSession.objects.unscoped().get(sandbox_id="sb-claim").user_attributed_at is not None

def test_desktop_claim_updates_run_and_open_session_compute_source(self):
run = self._run(state={"prewarmed": True, "await_user_message": True})
open_sandbox_session(run_id=run.id, sandbox_id="sb-code-claim", config=_config())

record_task_run_user_activity(run.id, self.team.id, compute_source=ComputeSource.POSTHOG_DESKTOP)

run.refresh_from_db()
session = SandboxSession.objects.unscoped().get(sandbox_id="sb-code-claim")
assert run.compute_source == ComputeSource.POSTHOG_DESKTOP
assert session.compute_source == ComputeSource.POSTHOG_DESKTOP

def test_later_desktop_activity_does_not_relabel_an_attributed_session(self):
run = self._run(state={"prewarmed": True, "await_user_message": True})
open_sandbox_session(run_id=run.id, sandbox_id="sb-other-claim", config=_config())
record_task_run_user_activity(run.id, self.team.id)

record_task_run_user_activity(run.id, self.team.id, compute_source=ComputeSource.POSTHOG_DESKTOP)

run.refresh_from_db()
session = SandboxSession.objects.unscoped().get(sandbox_id="sb-other-claim")
assert run.compute_source is None
assert session.compute_source is None


class TestSandboxUsageAggregation(SandboxUsageBase):
BEGIN = datetime(2026, 1, 2, tzinfo=UTC)
Expand Down
Loading
Loading