Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
19 commits
Select commit Hold shift + click to select a range
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
25 changes: 15 additions & 10 deletions backend/adapter_processor_v2/adapter_processor.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,13 @@
from django.core.exceptions import ObjectDoesNotExist
from platform_settings_v2.platform_auth_service import PlatformAuthenticationService
from tenant_account_v2.organization_member_service import OrganizationMemberService
from unstract.sdk1.adapters.adapterkit import Adapterkit
from unstract.sdk1.adapters.base import Adapter
from unstract.sdk1.adapters.x2text.constants import X2TextConstants
from unstract.sdk1.constants import AdapterTypes
from unstract.sdk1.embedding import EmbeddingCompat
from unstract.sdk1.exceptions import SdkError
from unstract.sdk1.llm import LLM

from adapter_processor_v2.constants import AdapterKeys, AllowedDomains
from adapter_processor_v2.exceptions import (
Expand All @@ -16,13 +23,6 @@
InValidAdapterId,
TestAdapterError,
)
from unstract.sdk1.adapters.adapterkit import Adapterkit
from unstract.sdk1.adapters.base import Adapter
from unstract.sdk1.adapters.x2text.constants import X2TextConstants
from unstract.sdk1.constants import AdapterTypes
from unstract.sdk1.embedding import EmbeddingCompat
from unstract.sdk1.exceptions import SdkError
from unstract.sdk1.llm import LLM

from .models import AdapterInstance, UserDefaultAdapter

Expand All @@ -46,6 +46,9 @@ def get_json_schema(adapter_id: str) -> dict[str, Any]:
schema_details[AdapterKeys.JSON_SCHEMA] = json.loads(
updated_adapters[0].get(AdapterKeys.JSON_SCHEMA)
)
for key in ("oauth", "oauth_provider", "python_social_auth_backend"):
if key in updated_adapters[0]:
schema_details[key] = updated_adapters[0][key]
else:
logger.error(f"Invalid adapter Id : {adapter_id} while fetching JSON Schema")
raise InValidAdapterId()
Expand All @@ -68,16 +71,18 @@ def get_all_supported_adapters(user_email: str, type: str) -> list[dict[Any, Any
if not is_special_user and adapter_id.startswith("noOp"):
continue

supported_adapters.append(
{
adapter_details = {
AdapterKeys.ID: adapter_id,
AdapterKeys.NAME: each_adapter.get(AdapterKeys.NAME),
AdapterKeys.DESCRIPTION: each_adapter.get(AdapterKeys.DESCRIPTION),
AdapterKeys.ICON: each_adapter.get(AdapterKeys.ICON),
AdapterKeys.ADAPTER_TYPE: each_adapter.get(AdapterKeys.ADAPTER_TYPE),
AdapterKeys.DOC_URL: each_adapter.get(AdapterKeys.DOC_URL, ""),
}
)
for key in ("oauth", "oauth_provider", "python_social_auth_backend"):
if key in each_adapter:
adapter_details[key] = each_adapter[key]
supported_adapters.append(adapter_details)
return supported_adapters

@staticmethod
Expand Down
15 changes: 11 additions & 4 deletions backend/adapter_processor_v2/serializers.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,9 @@
from typing import Any

from account_v2.serializer import UserSerializer
from backend.constants import FieldLengthConstants as FLC
from backend.serializers import AuditSerializer
from connector_auth_v2.openai_oauth import redact_openai_oauth_metadata
from cryptography.fernet import Fernet
from django.conf import settings
from rest_framework import serializers
Expand All @@ -10,14 +13,13 @@
serialize_group_refs,
serialize_owner_refs,
)
from unstract.sdk1.auth.openai_oauth import is_openai_oauth_adapter
from unstract.sdk1.constants import AdapterTypes
from unstract.sdk1.constants import Common as common
from utils.input_sanitizer import validate_name_field, validate_no_html_tags

from adapter_processor_v2.adapter_processor import AdapterProcessor
from adapter_processor_v2.constants import AdapterKeys
from backend.constants import FieldLengthConstants as FLC
from backend.serializers import AuditSerializer
from unstract.sdk1.constants import AdapterTypes
from unstract.sdk1.constants import Common as common

from .models import AdapterInstance, UserDefaultAdapter

Expand Down Expand Up @@ -87,6 +89,11 @@ def to_representation(self, instance: AdapterInstance) -> dict[str, str]:
rep.pop(AdapterKeys.ADAPTER_METADATA_B)
adapter_metadata = instance.metadata

if is_openai_oauth_adapter(instance.adapter_id):
# OAuth tokens and account identity stay encrypted in the database
# and are never returned to the browser after the initial login.
adapter_metadata = redact_openai_oauth_metadata(adapter_metadata)

# Hide unstract_key when use_platform_provided_unstract_key is True
if (
adapter_metadata.get("use_platform_provided_unstract_key") is True
Expand Down
137 changes: 129 additions & 8 deletions backend/adapter_processor_v2/views.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@
from typing import Any

from account_v2.models import User
from connector_auth_v2.openai_oauth import OpenAIOAuthService
from django.db import IntegrityError
from django.db.models import ProtectedError, QuerySet
from django.http import HttpRequest
Expand All @@ -19,14 +20,19 @@
from plugins import get_plugin
from rest_framework import status
from rest_framework.decorators import action
from rest_framework.exceptions import PermissionDenied
from rest_framework.exceptions import PermissionDenied, ValidationError
from rest_framework.request import Request
from rest_framework.response import Response
from rest_framework.serializers import ModelSerializer
from rest_framework.versioning import URLPathVersioning
from rest_framework.viewsets import GenericViewSet, ModelViewSet
from tenant_account_v2.organization_member_service import OrganizationMemberService
from tool_instance_v2.models import ToolInstance
from unstract.sdk1.auth.openai_oauth import (
OPENAI_OAUTH_PRIVATE_FIELDS,
OpenAIOAuthError,
is_openai_oauth_adapter,
)
from utils.filtering import FilterHelper
from utils.pagination import OptionalPagination
from utils.user_context import UserContext
Expand Down Expand Up @@ -60,6 +66,88 @@
logger = logging.getLogger(__name__)


def _prepare_openai_oauth_payload(
request: Request,
payload: Any,
existing_metadata: dict[str, Any] | None = None,
) -> tuple[Any, str | None]:
"""Inject credentials from one owned OAuth login into an adapter payload."""
adapter_id = payload.get(AdapterKeys.ADAPTER_ID)
if not is_openai_oauth_adapter(adapter_id):
return payload, None

submitted_metadata = payload.get(AdapterKeys.ADAPTER_METADATA)
if submitted_metadata is None and existing_metadata is not None:
adapter_metadata = dict(existing_metadata)
elif isinstance(submitted_metadata, dict):
adapter_metadata = dict(submitted_metadata)
else:
adapter_metadata = {}

# Tokens/account identity are always sourced from the server-side login
# session or the already-encrypted row. Never trust values sent in JSON.
for key in OPENAI_OAUTH_PRIVATE_FIELDS:
adapter_metadata.pop(key, None)
adapter_metadata.pop("oauth_authenticated", None)
adapter_metadata.pop("oauth_account_label", None)

oauth_key = request.query_params.get("oauth-key")
if oauth_key:
try:
credentials = OpenAIOAuthService.credentials_for_request(oauth_key, request)
except OpenAIOAuthError as exc:
raise ValidationError({"oauth-key": str(exc)}) from exc
adapter_metadata.update(credentials)
elif existing_metadata is not None:
for key in OPENAI_OAUTH_PRIVATE_FIELDS:
if key in existing_metadata:
adapter_metadata[key] = existing_metadata[key]
else:
raise ValidationError(
{"oauth-key": "OpenAI OAuth authentication is required for this adapter."}
)

payload[AdapterKeys.ADAPTER_METADATA] = adapter_metadata
return payload, oauth_key


def _saved_openai_oauth_metadata(request: Request) -> dict[str, Any] | None:
"""Load one owned adapter's encrypted OAuth metadata for a test request."""
adapter_instance_id = request.query_params.get("adapter-instance-id")
if not adapter_instance_id:
return None
try:
adapter_uuid = uuid.UUID(adapter_instance_id)
except (AttributeError, TypeError, ValueError) as exc:
raise ValidationError(
{"adapter-instance-id": "OpenAI OAuth adapter was not found"}
) from exc

adapter = (
AdapterInstance.objects.for_user(request.user)
.filter(pk=adapter_uuid)
.first()
)
if adapter is None or not is_openai_oauth_adapter(adapter.adapter_id):
raise ValidationError(
{"adapter-instance-id": "OpenAI OAuth adapter was not found"}
)
metadata = adapter.metadata
return metadata if isinstance(metadata, dict) else None


def _consume_openai_oauth_key(cache_key: str | None, request: Request) -> None:
"""Best-effort cleanup after credentials are durably stored."""
if not cache_key:
return
try:
OpenAIOAuthService.consume(cache_key, request)
except OpenAIOAuthError:
# The adapter row is already the durable credential store. A cache
# expiry/race must not turn a successful create/update into a 500.
logger.warning("Could not consume OpenAI OAuth hand-off session")


class DefaultAdapterViewSet(ModelViewSet):
versioning_class = URLPathVersioning
serializer_class = DefaultAdapterSerializer
Expand Down Expand Up @@ -127,10 +215,18 @@ def get_adapter_schema(

def test(self, request: Request) -> Response:
"""Tests the connector against the credentials passed."""
serializer: AdapterInstanceSerializer = self.get_serializer(data=request.data)
payload = request.data.copy()
existing_metadata = None
if is_openai_oauth_adapter(payload.get(AdapterKeys.ADAPTER_ID)):
existing_metadata = _saved_openai_oauth_metadata(request)
payload, _ = _prepare_openai_oauth_payload(
request, payload, existing_metadata=existing_metadata
)
serializer: AdapterInstanceSerializer = self.get_serializer(data=payload)
serializer.is_valid(raise_exception=True)
adapter_id = serializer.validated_data.get(AdapterKeys.ADAPTER_ID)
adapter_metadata = serializer.validated_data.get(AdapterKeys.ADAPTER_METADATA)
adapter_metadata = dict(adapter_metadata or {})
adapter_metadata[AdapterKeys.ADAPTER_TYPE] = serializer.validated_data.get(
AdapterKeys.ADAPTER_TYPE
)
Expand Down Expand Up @@ -237,10 +333,12 @@ def _enforce_llm_creation_restriction(request: Any, adapter_type: str) -> None:
)

def create(self, request: Any) -> Response:
serializer = self.get_serializer(data=request.data)
payload = request.data.copy()
payload, oauth_key = _prepare_openai_oauth_payload(request, payload)
serializer = self.get_serializer(data=payload)

use_platform_unstract_key = False
adapter_metadata = request.data.get(AdapterKeys.ADAPTER_METADATA)
adapter_metadata = payload.get(AdapterKeys.ADAPTER_METADATA)
if adapter_metadata and adapter_metadata.get(
AdapterKeys.PLATFORM_PROVIDED_UNSTRACT_KEY, False
):
Expand Down Expand Up @@ -312,6 +410,9 @@ def create(self, request: Any) -> Response:

user_default_adapter.save()

# The encrypted adapter row is now the durable credential store.
_consume_openai_oauth_key(oauth_key, request)

except IntegrityError:
raise DuplicateAdapterNameError(
name=serializer.validated_data.get(AdapterKeys.ADAPTER_NAME)
Expand Down Expand Up @@ -410,7 +511,10 @@ def partial_update(
adapter = self.get_object()
before = self.snapshot_share_axes(adapter)

response = super().partial_update(request, *args, **kwargs)
if is_openai_oauth_adapter(adapter.adapter_id):
response = self._update_openai_oauth(request, adapter, partial=True)
else:
response = super().partial_update(request, *args, **kwargs)
if response.status_code == 200 and notification_plugin:
self._notify_shared_users(adapter, before, request.data, request.user)
return response
Expand Down Expand Up @@ -555,6 +659,13 @@ def list_of_shared_users(self, request: HttpRequest, pk: Any = None) -> Response
def update(
self, request: Request, *args: tuple[Any], **kwargs: dict[str, Any]
) -> Response:
# OAuth adapters carry their credentials in the encrypted metadata row;
# inject a newly completed login or preserve the existing account when
# a metadata-only edit does not include a new login session.
adapter = self.get_object()
if is_openai_oauth_adapter(adapter.adapter_id):
return self._update_openai_oauth(request, adapter, partial=False)

# Check if adapter metadata is being updated and contains the platform key flag
use_platform_unstract_key = False
adapter_metadata = request.data.get(AdapterKeys.ADAPTER_METADATA)
Expand All @@ -565,9 +676,6 @@ def update(
use_platform_unstract_key = True
logger.error(f"Platform key flag detected: {use_platform_unstract_key}")

# Get the adapter instance for update
adapter = self.get_object()

if use_platform_unstract_key:
logger.error("Processing adapter with platform key")
serializer = self.get_serializer(adapter, data=request.data, partial=True)
Expand Down Expand Up @@ -597,6 +705,19 @@ def update(
# For non-platform-key cases, use the default update behavior
return super().update(request, *args, **kwargs)

def _update_openai_oauth(
self, request: Request, adapter: AdapterInstance, *, partial: bool
) -> Response:
payload = request.data.copy()
payload, oauth_key = _prepare_openai_oauth_payload(
request, payload, existing_metadata=adapter.metadata
)
serializer = self.get_serializer(adapter, data=payload, partial=partial)
serializer.is_valid(raise_exception=True)
serializer.save()
_consume_openai_oauth_key(oauth_key, request)
return Response(serializer.data)

@action(detail=True, methods=["get"])
def adapter_info(self, request: HttpRequest, pk: uuid) -> Response:
adapter = self.get_object()
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,68 @@
import uuid

import django.db.models.deletion
from django.conf import settings
from django.db import migrations, models


class Migration(migrations.Migration):
dependencies = [
migrations.swappable_dependency(settings.AUTH_USER_MODEL),
("connector_auth_v2", "0001_initial"),
]

operations = [
migrations.CreateModel(
name="OpenAIOAuthCredential",
fields=[
(
"id",
models.UUIDField(
default=uuid.uuid4,
editable=False,
primary_key=True,
serialize=False,
),
),
(
"organization_id",
models.CharField(max_length=64),
),
("account_id", models.CharField(max_length=255)),
(
"account_label",
models.CharField(blank=True, default="", max_length=255),
),
("encrypted_credentials", models.TextField()),
("created_at", models.DateTimeField(auto_now_add=True)),
("modified_at", models.DateTimeField(auto_now=True)),
(
"user",
models.ForeignKey(
on_delete=django.db.models.deletion.CASCADE,
related_name="openai_oauth_credentials",
to=settings.AUTH_USER_MODEL,
),
),
],
options={
"verbose_name": "OpenAI OAuth credential",
"verbose_name_plural": "OpenAI OAuth credentials",
"db_table": "openai_oauth_credential",
},
),
migrations.AddConstraint(
model_name="openaioauthcredential",
constraint=models.UniqueConstraint(
fields=("user", "organization_id", "account_id"),
name="unique_openai_oauth_user_org_account",
),
),
migrations.AddIndex(
model_name="openaioauthcredential",
index=models.Index(
fields=("user", "organization_id", "-modified_at"),
name="openai_oauth_user_org_mod_idx",
),
),
]
Loading