Skip to content
Draft
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
71 changes: 67 additions & 4 deletions apps/backend/fastapi/main.py
Original file line number Diff line number Diff line change
@@ -1,16 +1,21 @@
import math
import os
import secrets
import sys
from contextlib import asynccontextmanager
from datetime import date, datetime
from decimal import Decimal
from pathlib import Path
from typing import Annotated

import numpy as np
import pandas as pd
import sqlglot
from sqlglot import exp
from sqlglot.errors import ParseError
import uvicorn
from dotenv import load_dotenv
from fastapi import FastAPI, HTTPException
from fastapi import Depends, FastAPI, Header, HTTPException
from fastapi.middleware.cors import CORSMiddleware
from pydantic import BaseModel

Expand All @@ -23,6 +28,8 @@
from nao_core.context import get_context_provider

port = int(os.environ.get("PORT", 8005))
INTERNAL_AUTH_HEADER = "X-Nao-Internal-Secret"
MIN_INTERNAL_SECRET_LENGTH = 20

# Global scheduler instance
scheduler = None
Expand Down Expand Up @@ -97,6 +104,7 @@ class ExecuteSQLRequest(BaseModel):
database_id: str | None = None
env_vars: dict[str, str] | None = None
azure_access_token: str | None = None
dangerously_write_permission_enabled: bool = False


class ExecuteSQLResponse(BaseModel):
Expand Down Expand Up @@ -161,6 +169,45 @@ def _convert_value(v: object):
return v


def _require_internal_auth(
provided_secret: Annotated[
str | None,
Header(alias=INTERNAL_AUTH_HEADER),
] = None,
):
expected_secret = os.environ.get("BETTER_AUTH_SECRET")
if not expected_secret or len(expected_secret) < MIN_INTERNAL_SECRET_LENGTH:
raise HTTPException(
status_code=503,
detail="Internal API authentication is not configured",
)

provided_bytes = provided_secret.encode() if provided_secret is not None else b""
if not secrets.compare_digest(provided_bytes, expected_secret.encode()):
raise HTTPException(status_code=401, detail="Invalid internal API credentials")


def _is_read_only_sql(sql: str) -> bool:
try:
statements = sqlglot.parse(sql, error_level=sqlglot.ErrorLevel.RAISE)
except (ParseError, ValueError):
return False

if not statements:
return False

forbidden_expression_types = (exp.DDL, exp.DML, exp.Into, exp.Lock)
return all(
statement is not None
and isinstance(statement, exp.Query)
and not any(
isinstance(expression, forbidden_expression_types)
for expression in statement.walk()
)
for statement in statements
)


# =============================================================================
# API Endpoints
# =============================================================================
Expand All @@ -187,7 +234,11 @@ async def health_check():
)


@app.post("/api/refresh", response_model=RefreshResponse)
@app.post(
"/api/refresh",
response_model=RefreshResponse,
dependencies=[Depends(_require_internal_auth)],
)
async def refresh_context():
"""Trigger a context refresh (git pull if using git source).

Expand Down Expand Up @@ -219,9 +270,21 @@ async def refresh_context():
)


@app.post("/execute_sql", response_model=ExecuteSQLResponse)
@app.post(
"/execute_sql",
response_model=ExecuteSQLResponse,
dependencies=[Depends(_require_internal_auth)],
)
async def execute_sql(request: ExecuteSQLRequest):
try:
if not request.dangerously_write_permission_enabled and not _is_read_only_sql(
request.sql
):
raise HTTPException(
status_code=403,
detail="Write SQL operations are disabled",
)

project_path = Path(request.nao_project_folder)
config = NaoConfig.try_load(
project_path,
Expand Down Expand Up @@ -300,4 +363,4 @@ async def execute_sql(request: ExecuteSQLRequest):


if __name__ == "__main__":
uvicorn.run("main:app", host="0.0.0.0", port=port, reload=True)
uvicorn.run("main:app", host="127.0.0.1", port=port, reload=True)
118 changes: 118 additions & 0 deletions apps/backend/fastapi/test_main.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,4 @@
import os
import tempfile
from pathlib import Path

Expand All @@ -6,6 +7,11 @@
from fastapi.testclient import TestClient
from main import app

TEST_INTERNAL_SECRET = "test-internal-secret-at-least-20-characters"
os.environ["BETTER_AUTH_SECRET"] = TEST_INTERNAL_SECRET

AUTH_HEADERS = {"X-Nao-Internal-Secret": TEST_INTERNAL_SECRET}


def assert_sql_result(
data: dict, *, row_count: int, columns: list[str], expected_data: list[dict]
Expand Down Expand Up @@ -43,6 +49,7 @@ def test_execute_sql_simple_duckdb(duckdb_project_folder):

response = client.post(
"/execute_sql",
headers=AUTH_HEADERS,
json={
"sql": "SELECT 1 AS id, 'hello' AS message",
"nao_project_folder": duckdb_project_folder,
Expand All @@ -64,6 +71,7 @@ def test_execute_sql_with_cte_duckdb(duckdb_project_folder):

response = client.post(
"/execute_sql",
headers=AUTH_HEADERS,
json={
"sql": "WITH test AS (SELECT 1 AS id, 'hello' AS message) SELECT * FROM test",
"nao_project_folder": duckdb_project_folder,
Expand All @@ -79,6 +87,114 @@ def test_execute_sql_with_cte_duckdb(duckdb_project_folder):
)


def test_health_does_not_require_internal_auth():
response = TestClient(app).get("/health")

assert response.status_code == 200


@pytest.mark.parametrize(
"headers",
[
{},
{"X-Nao-Internal-Secret": "incorrect-secret-at-least-20-characters"},
],
)
def test_execute_sql_requires_internal_auth(duckdb_project_folder, headers):
response = TestClient(app).post(
"/execute_sql",
headers=headers,
json={
"sql": "SELECT 1",
"nao_project_folder": duckdb_project_folder,
},
)

assert response.status_code == 401


def test_execute_sql_fails_closed_without_configured_secret(
duckdb_project_folder,
monkeypatch,
):
monkeypatch.delenv("BETTER_AUTH_SECRET")

response = TestClient(app).post(
"/execute_sql",
headers=AUTH_HEADERS,
json={
"sql": "SELECT 1",
"nao_project_folder": duckdb_project_folder,
},
)

assert response.status_code == 503


def test_refresh_requires_internal_auth():
response = TestClient(app).post("/api/refresh")

assert response.status_code == 401


@pytest.mark.parametrize(
"sql",
[
"DELETE FROM users",
"SELECT 1; DROP TABLE users",
"WITH deleted AS (DELETE FROM users RETURNING *) SELECT * FROM deleted",
"SELECT * INTO copied_users FROM users",
"SELECT * FROM users FOR UPDATE",
"SELECT (",
],
)
def test_execute_sql_rejects_non_read_only_sql(duckdb_project_folder, sql):
response = TestClient(app).post(
"/execute_sql",
headers=AUTH_HEADERS,
json={
"sql": sql,
"nao_project_folder": duckdb_project_folder,
},
)

assert response.status_code == 403
assert response.json()["detail"] == "Write SQL operations are disabled"


def test_execute_sql_allows_authenticated_write_permission(monkeypatch):
class FakeDatabase:
name = "test"
type = "duckdb"
auth_mode = None

def execute_sql(self, sql):
import pandas as pd

assert sql == "DELETE FROM users"
return pd.DataFrame()

class FakeConfig:
databases = [FakeDatabase()]

monkeypatch.setattr(
"main.NaoConfig.try_load",
lambda *args, **kwargs: FakeConfig(),
)

response = TestClient(app).post(
"/execute_sql",
headers=AUTH_HEADERS,
json={
"sql": "DELETE FROM users",
"nao_project_folder": "/tmp/test-project",
"dangerously_write_permission_enabled": True,
},
)

assert response.status_code == 200


# BigQuery tests (requires SSO authentication)


Expand Down Expand Up @@ -109,6 +225,7 @@ def test_execute_sql_simple_bigquery(bigquery_project_folder):

response = client.post(
"/execute_sql",
headers=AUTH_HEADERS,
json={
"sql": "SELECT 1 AS id, 'hello' AS message",
"nao_project_folder": bigquery_project_folder,
Expand Down Expand Up @@ -139,6 +256,7 @@ def test_execute_sql_with_cte_bigquery(bigquery_project_folder):

response = client.post(
"/execute_sql",
headers=AUTH_HEADERS,
json={
"sql": cte_sql,
"nao_project_folder": bigquery_project_folder,
Expand Down
2 changes: 2 additions & 0 deletions apps/backend/src/agents/tools/execute-sql.ts
Original file line number Diff line number Diff line change
Expand Up @@ -30,10 +30,12 @@ export async function executeQuery(
method: 'POST',
headers: {
'Content-Type': 'application/json',
'X-Nao-Internal-Secret': env.BETTER_AUTH_SECRET,
},
body: JSON.stringify({
sql: sql_query,
nao_project_folder: naoProjectFolder,
dangerously_write_permission_enabled: writePermEnabled,
...(database_id && { database_id }),
...(Object.keys(envVars).length > 0 && { env_vars: envVars }),
...(context.azureAccessToken && { azure_access_token: context.azureAccessToken }),
Expand Down
12 changes: 7 additions & 5 deletions apps/backend/src/services/context-explorer.service.ts
Original file line number Diff line number Diff line change
Expand Up @@ -9,7 +9,7 @@ export async function getFileTree(projectFolder: string): Promise<FileTreeEntry[
}

export async function readFileContent(filePath: string, projectFolder: string): Promise<string> {
const realPath = resolveAndValidatePath(filePath, projectFolder);
const realPath = await resolveAndValidatePath(filePath, projectFolder);
const stat = await fs.stat(realPath);

const MAX_FILE_SIZE = 1024 * 1024; // 1 MB
Expand Down Expand Up @@ -59,14 +59,16 @@ async function readDirectoryRecursive(dirPath: string, projectFolder: string): P
return entries;
}

function resolveAndValidatePath(virtualPath: string, projectFolder: string): string {
async function resolveAndValidatePath(virtualPath: string, projectFolder: string): Promise<string> {
const relativePath = virtualPath.startsWith('/') ? virtualPath.slice(1) : virtualPath;
const resolvedPath = path.resolve(projectFolder, relativePath);
const realProjectFolder = await fs.realpath(projectFolder);
const resolvedPath = path.resolve(realProjectFolder, relativePath);
const realPath = await fs.realpath(resolvedPath);

const withinFolder = resolvedPath === projectFolder || resolvedPath.startsWith(projectFolder + path.sep);
const withinFolder = realPath === realProjectFolder || realPath.startsWith(realProjectFolder + path.sep);
if (!withinFolder) {
throw new Error(`Access denied: path is outside the project folder`);
}

return resolvedPath;
return realPath;
}
5 changes: 4 additions & 1 deletion apps/backend/src/services/live-story.ts
Original file line number Diff line number Diff line change
Expand Up @@ -130,7 +130,10 @@ async function executeRawSql(
): Promise<{ data: unknown[]; columns: string[] }> {
const response = await fetch(`http://localhost:${env.FASTAPI_PORT}/execute_sql`, {
method: 'POST',
headers: { 'Content-Type': 'application/json' },
headers: {
'Content-Type': 'application/json',
'X-Nao-Internal-Secret': env.BETTER_AUTH_SECRET,
},
body: JSON.stringify({
sql: sqlQuery,
nao_project_folder: projectFolder,
Expand Down
4 changes: 4 additions & 0 deletions apps/backend/src/trpc/account.routes.ts
Original file line number Diff line number Diff line change
Expand Up @@ -2,11 +2,13 @@ import { TRPCError } from '@trpc/server';
import { hashPassword } from 'better-auth/crypto';
import { z } from 'zod/v4';

import { isCloud } from '../env';
import * as accountQueries from '../queries/account.queries';
import * as projectQueries from '../queries/project.queries';
import * as userQueries from '../queries/user.queries';
import { emailService } from '../services/email';
import { buildResetPasswordEmail } from '../utils/email-builders';
import { assertAdminPasswordResetAllowed } from '../utils/password-reset';
import { regexPassword } from '../utils/utils';
import { adminProtectedProcedure, protectedProcedure } from './trpc';

Expand All @@ -18,6 +20,8 @@ export const accountRoutes = {
}),
)
.mutation(async ({ input, ctx }) => {
assertAdminPasswordResetAllowed(isCloud);

const account = await accountQueries.getAccountById(input.userId);
if (!account || !account.password) {
throw new TRPCError({
Expand Down
3 changes: 3 additions & 0 deletions apps/backend/src/trpc/organization.routes.ts
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,7 @@ import { emailService } from '../services/email';
import { addTeamMember } from '../services/team-member';
import { ORG_ROLES } from '../types/organization';
import { buildResetPasswordEmail, buildUserAddedEmail } from '../utils/email-builders';
import { assertAdminPasswordResetAllowed } from '../utils/password-reset';
import { isPublicEmailDomain, normalizeEmailDomains } from '../utils/utils';
import { protectedProcedure } from './trpc';

Expand Down Expand Up @@ -138,6 +139,8 @@ export const organizationRoutes = {
resetMemberPassword: orgAdminOnlyProcedure
.input(z.object({ userId: z.string() }))
.mutation(async ({ input, ctx }) => {
assertAdminPasswordResetAllowed(isCloud);

const memberRole = await orgQueries.getUserRoleInOrg(ctx.org.id, input.userId);
if (!memberRole) {
throw new TRPCError({ code: 'FORBIDDEN', message: 'User is not a member of this organization.' });
Expand Down
Loading
Loading