Rewrite identity and access management

Signed-off-by: Bryan Frimin <bryan@getprobo.com>
This commit is contained in:
Bryan Frimin
2025-12-03 19:23:15 +01:00
parent 4ed3f5a067
commit 74fc3b8cd1
201 changed files with 32895 additions and 23649 deletions

2
.gitattributes vendored
View File

@@ -1 +1 @@
pkg/server/api/console/v1/v1_resolver.go linguist-generated=false pkg/server/api/console/v1/v1_resolver.go linguist-detectable

View File

@@ -18,6 +18,7 @@
"@probo/hooks": "1.0.0", "@probo/hooks": "1.0.0",
"@probo/i18n": "1.0.0", "@probo/i18n": "1.0.0",
"@probo/react-lazy": "1.0.0", "@probo/react-lazy": "1.0.0",
"@probo/relay": "^1.0.0",
"@probo/routes": "^1.0.0", "@probo/routes": "^1.0.0",
"@probo/ui": "1.0.0", "@probo/ui": "1.0.0",
"@probo/vendors": "0.0.1", "@probo/vendors": "0.0.1",

View File

@@ -2,7 +2,7 @@ import { useLocation, useRouteError } from "react-router";
import { IconPageCross } from "@probo/ui"; import { IconPageCross } from "@probo/ui";
import { useTranslate } from "@probo/i18n"; import { useTranslate } from "@probo/i18n";
import { useEffect, useRef } from "react"; import { useEffect, useRef } from "react";
import { AuthenticationRequiredError } from "/providers/RelayProviders"; import { AuthenticationRequiredError } from "@probo/relay";
const classNames = { const classNames = {
wrapper: "py-10 text-center space-y-2 ", wrapper: "py-10 text-center space-y-2 ",

View File

@@ -0,0 +1,18 @@
import { Environment, Network, RecordSource, Store } from "relay-runtime";
import { makeFetchQuery } from "@probo/relay";
const source = new RecordSource();
const store = new Store(source, {
queryCacheExpirationTime: 1 * 60 * 1000,
gcReleaseBufferSize: 20,
});
export const consoleEnvironment = new Environment({
network: Network.create(makeFetchQuery("/api/console/v1/query")),
store,
});
export const connectEnvironment = new Environment({
network: Network.create(makeFetchQuery("/api/connect/v1/query")),
store,
});

View File

@@ -32,7 +32,6 @@ import type { EmployeeLayoutQuery as EmployeeLayoutQueryType } from "./__generat
import { Suspense, useState, useEffect, useMemo, use } from "react"; import { Suspense, useState, useEffect, useMemo, use } from "react";
import { ErrorBoundary } from "react-error-boundary"; import { ErrorBoundary } from "react-error-boundary";
import { PageError } from "/components/PageError"; import { PageError } from "/components/PageError";
import { buildEndpoint } from "/providers/RelayProviders";
import { PermissionsProvider } from "/providers/PermissionsProvider"; import { PermissionsProvider } from "/providers/PermissionsProvider";
import { PermissionsContext } from "/providers/PermissionsContext"; import { PermissionsContext } from "/providers/PermissionsContext";
@@ -71,11 +70,7 @@ export function EmployeeLayout() {
); );
} }
function EmployeeLayoutContent({ function EmployeeLayoutContent({ organizationId }: { organizationId: string }) {
organizationId,
}: {
organizationId: string;
}) {
const data = useLazyLoadQuery<EmployeeLayoutQueryType>(EmployeeLayoutQuery, { const data = useLazyLoadQuery<EmployeeLayoutQueryType>(EmployeeLayoutQuery, {
organizationId, organizationId,
}); });
@@ -352,7 +347,7 @@ function UserDropdown({ organizationId }: { organizationId: string }) {
) => { ) => {
e.preventDefault(); e.preventDefault();
fetch(buildEndpoint("/connect/logout"), { fetch("/connect/logout", {
method: "DELETE", method: "DELETE",
headers: { headers: {
"Content-Type": "application/json", "Content-Type": "application/json",

View File

@@ -50,7 +50,6 @@ import { graphql } from "relay-runtime";
import type { MainLayoutQuery as MainLayoutQueryType } from "./__generated__/MainLayoutQuery.graphql"; import type { MainLayoutQuery as MainLayoutQueryType } from "./__generated__/MainLayoutQuery.graphql";
import { PageError } from "/components/PageError"; import { PageError } from "/components/PageError";
import { PermissionsProvider } from "/providers/PermissionsProvider"; import { PermissionsProvider } from "/providers/PermissionsProvider";
import { buildEndpoint } from "/providers/RelayProviders";
import { PermissionsContext } from "/providers/PermissionsContext"; import { PermissionsContext } from "/providers/PermissionsContext";
const MainLayoutQuery = graphql` const MainLayoutQuery = graphql`
@@ -283,7 +282,7 @@ function UserDropdown({ organizationId }: { organizationId: string }) {
) => { ) => {
e.preventDefault(); e.preventDefault();
fetch(buildEndpoint("/connect/logout"), { fetch("/connect/logout", {
method: "DELETE", method: "DELETE",
headers: { headers: {
"Content-Type": "application/json", "Content-Type": "application/json",
@@ -371,7 +370,6 @@ interface InvitationsResponse {
invitations: Invitation[]; invitations: Invitation[];
} }
function OrganizationSelector({ function OrganizationSelector({
currentOrganization, currentOrganization,
}: { }: {

View File

@@ -1,7 +1,6 @@
import { Outlet } from "react-router"; import { Outlet } from "react-router";
import { Logo, Button, IconArrowBoxLeft, useToast } from "@probo/ui"; import { Logo, Button, IconArrowBoxLeft, useToast } from "@probo/ui";
import { useTranslate } from "@probo/i18n"; import { useTranslate } from "@probo/i18n";
import { buildEndpoint } from "/providers/RelayProviders";
import type { ReactNode } from "react"; import type { ReactNode } from "react";
type Props = { type Props = {
@@ -11,18 +10,23 @@ type Props = {
isAuthenticated?: boolean; isAuthenticated?: boolean;
}; };
export function PublicTrustCenterLayout({ organizationName, organizationLogo, children, isAuthenticated }: Props) { export function PublicTrustCenterLayout({
organizationName,
organizationLogo,
children,
isAuthenticated,
}: Props) {
const { __ } = useTranslate(); const { __ } = useTranslate();
const { toast } = useToast(); const { toast } = useToast();
const handleLogout = async () => { const handleLogout = async () => {
try { try {
const response = await fetch(buildEndpoint('/api/trust/v1/auth/logout'), { const response = await fetch("/api/trust/v1/auth/logout", {
method: 'DELETE', method: "DELETE",
headers: { headers: {
'Content-Type': 'application/json', "Content-Type": "application/json",
}, },
credentials: 'include', credentials: "include",
}); });
if (!response.ok) { if (!response.ok) {

View File

@@ -1,7 +1,6 @@
import { createRoot } from "react-dom/client"; import { createRoot } from "react-dom/client";
import "./index.css"; import "./index.css";
import { App } from "./App"; import { App } from "./App";
import { RelayProvider } from "./providers/RelayProviders";
import { TranslatorProvider } from "./providers/TranslatorProvider"; import { TranslatorProvider } from "./providers/TranslatorProvider";
import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; import { QueryClient, QueryClientProvider } from "@tanstack/react-query";
@@ -15,10 +14,8 @@ const queryClient = new QueryClient({
createRoot(document.getElementById("root")!).render( createRoot(document.getElementById("root")!).render(
<QueryClientProvider client={queryClient}> <QueryClientProvider client={queryClient}>
<RelayProvider> <TranslatorProvider>
<TranslatorProvider> <App />
<App /> </TranslatorProvider>
</TranslatorProvider> </QueryClientProvider>
</RelayProvider>
</QueryClientProvider>,
); );

View File

@@ -31,7 +31,7 @@ import { formatDate } from "@probo/helpers";
import { useFormWithSchema } from "/hooks/useFormWithSchema"; import { useFormWithSchema } from "/hooks/useFormWithSchema";
import { Controller } from "react-hook-form"; import { Controller } from "react-hook-form";
import { z } from "zod"; import { z } from "zod";
import { UnAuthenticatedError } from "/providers/RelayProviders"; import { UnAuthenticatedError } from "@probo/relay";
interface APIKey { interface APIKey {
id: string; id: string;
@@ -56,10 +56,14 @@ interface Organization {
const createSchema = z.object({ const createSchema = z.object({
name: z.string().min(1, "Name is required"), name: z.string().min(1, "Name is required"),
expiresIn: z.enum(["1month", "3months", "6months", "1year"]), expiresIn: z.enum(["1month", "3months", "6months", "1year"]),
organizations: z.array(z.object({ organizations: z
organizationId: z.string(), .array(
role: z.string(), z.object({
})).min(1, "At least one organization is required"), organizationId: z.string(),
role: z.string(),
})
)
.min(1, "At least one organization is required"),
}); });
type CreateFormData = z.infer<typeof createSchema>; type CreateFormData = z.infer<typeof createSchema>;
@@ -79,19 +83,24 @@ export default function APIKeysPage() {
const [isCreating, setIsCreating] = useState(false); const [isCreating, setIsCreating] = useState(false);
const [isUpdating, setIsUpdating] = useState(false); const [isUpdating, setIsUpdating] = useState(false);
const [isDeleting, setIsDeleting] = useState(false); const [isDeleting, setIsDeleting] = useState(false);
const [selectedOrganizations, setSelectedOrganizations] = useState<string[]>([]); const [selectedOrganizations, setSelectedOrganizations] = useState<string[]>(
const [organizationRoles, setOrganizationRoles] = useState<Record<string, string>>({}); []
);
const [organizationRoles, setOrganizationRoles] = useState<
Record<string, string>
>({});
const [editingAPIKey, setEditingAPIKey] = useState<APIKey | null>(null); const [editingAPIKey, setEditingAPIKey] = useState<APIKey | null>(null);
const [editingName, setEditingName] = useState<string>(""); const [editingName, setEditingName] = useState<string>("");
const [error, setError] = useState<Error | null>(null); const [error, setError] = useState<Error | null>(null);
const { formState, handleSubmit, reset, control, setValue, register } = useFormWithSchema(createSchema, { const { formState, handleSubmit, reset, control, setValue, register } =
defaultValues: { useFormWithSchema(createSchema, {
name: new Date().toISOString().split('T')[0], defaultValues: {
expiresIn: "1month", name: new Date().toISOString().split("T")[0],
organizations: [], expiresIn: "1month",
}, organizations: [],
}); },
});
if (error) { if (error) {
throw error; throw error;
@@ -99,14 +108,16 @@ export default function APIKeysPage() {
const fetchAPIKeys = async () => { const fetchAPIKeys = async () => {
try { try {
const response = await fetch('/connect/api-keys', { credentials: 'include' }); const response = await fetch("/connect/api-keys", {
credentials: "include",
});
if (!response.ok) { if (!response.ok) {
throw new Error('Failed to fetch API keys'); throw new Error("Failed to fetch API keys");
} }
const data: { apiKeys: APIKey[] } = await response.json(); const data: { apiKeys: APIKey[] } = await response.json();
setApiKeys(data.apiKeys); setApiKeys(data.apiKeys);
} catch (err) { } catch (err) {
console.error('Failed to fetch API keys:', err); console.error("Failed to fetch API keys:", err);
toast({ toast({
title: __("Error"), title: __("Error"),
description: __("Failed to load API keys"), description: __("Failed to load API keys"),
@@ -119,32 +130,40 @@ export default function APIKeysPage() {
const fetchData = async () => { const fetchData = async () => {
try { try {
const [apiKeysResponse, organizationsResponse] = await Promise.all([ const [apiKeysResponse, organizationsResponse] = await Promise.all([
fetch('/connect/api-keys', { credentials: 'include' }), fetch("/connect/api-keys", { credentials: "include" }),
fetch('/connect/organizations?role=OWNER', { credentials: 'include' }), fetch("/connect/organizations?role=OWNER", {
credentials: "include",
}),
]); ]);
if (apiKeysResponse.status === 401 || organizationsResponse.status === 401) { if (
apiKeysResponse.status === 401 ||
organizationsResponse.status === 401
) {
setError(new UnAuthenticatedError()); setError(new UnAuthenticatedError());
return; return;
} }
if (!apiKeysResponse.ok) { if (!apiKeysResponse.ok) {
throw new Error('Failed to fetch API keys'); throw new Error("Failed to fetch API keys");
} }
if (!organizationsResponse.ok) { if (!organizationsResponse.ok) {
throw new Error('Failed to fetch organizations'); throw new Error("Failed to fetch organizations");
} }
const apiKeysData: { apiKeys: APIKey[] } = await apiKeysResponse.json(); const apiKeysData: { apiKeys: APIKey[] } = await apiKeysResponse.json();
const orgsData: { organizations: Organization[] } = await organizationsResponse.json(); const orgsData: { organizations: Organization[] } =
await organizationsResponse.json();
const authenticatedOrgs = orgsData.organizations.filter(org => org.authStatus === "authenticated"); const authenticatedOrgs = orgsData.organizations.filter(
(org) => org.authStatus === "authenticated"
);
setApiKeys(apiKeysData.apiKeys); setApiKeys(apiKeysData.apiKeys);
setOrganizations(authenticatedOrgs); setOrganizations(authenticatedOrgs);
} catch (err) { } catch (err) {
console.error('Failed to fetch data:', err); console.error("Failed to fetch data:", err);
toast({ toast({
title: __("Error"), title: __("Error"),
description: __("Failed to load data"), description: __("Failed to load data"),
@@ -179,12 +198,12 @@ export default function APIKeysPage() {
setIsCreating(true); setIsCreating(true);
try { try {
const response = await fetch('/connect/api-keys', { const response = await fetch("/connect/api-keys", {
method: 'POST', method: "POST",
headers: { headers: {
'Content-Type': 'application/json', "Content-Type": "application/json",
}, },
credentials: 'include', credentials: "include",
body: JSON.stringify({ body: JSON.stringify({
name: formData.name, name: formData.name,
expiresAt: expiresAt.toISOString(), expiresAt: expiresAt.toISOString(),
@@ -193,7 +212,7 @@ export default function APIKeysPage() {
}); });
if (!response.ok) { if (!response.ok) {
throw new Error('Failed to create API key'); throw new Error("Failed to create API key");
} }
const data: { apiKey: APIKey; key: string } = await response.json(); const data: { apiKey: APIKey; key: string } = await response.json();
@@ -224,9 +243,9 @@ export default function APIKeysPage() {
const handleEdit = (apiKey: APIKey) => { const handleEdit = (apiKey: APIKey) => {
setEditingAPIKey(apiKey); setEditingAPIKey(apiKey);
setEditingName(apiKey.name); setEditingName(apiKey.name);
const orgIds = apiKey.organizations.map(org => org.organizationId); const orgIds = apiKey.organizations.map((org) => org.organizationId);
const roles: Record<string, string> = {}; const roles: Record<string, string> = {};
apiKey.organizations.forEach(org => { apiKey.organizations.forEach((org) => {
roles[org.organizationId] = org.role; roles[org.organizationId] = org.role;
}); });
setSelectedOrganizations(orgIds); setSelectedOrganizations(orgIds);
@@ -239,16 +258,16 @@ export default function APIKeysPage() {
setIsUpdating(true); setIsUpdating(true);
try { try {
const response = await fetch('/connect/api-keys', { const response = await fetch("/connect/api-keys", {
method: 'PUT', method: "PUT",
headers: { headers: {
'Content-Type': 'application/json', "Content-Type": "application/json",
}, },
credentials: 'include', credentials: "include",
body: JSON.stringify({ body: JSON.stringify({
id: editingAPIKey.id, id: editingAPIKey.id,
name: editingName, name: editingName,
organizations: selectedOrganizations.map(id => ({ organizations: selectedOrganizations.map((id) => ({
organizationId: id, organizationId: id,
role: organizationRoles[id] || "FULL", role: organizationRoles[id] || "FULL",
})), })),
@@ -256,7 +275,7 @@ export default function APIKeysPage() {
}); });
if (!response.ok) { if (!response.ok) {
throw new Error('Failed to update API key'); throw new Error("Failed to update API key");
} }
await fetchAPIKeys(); await fetchAPIKeys();
@@ -286,20 +305,20 @@ export default function APIKeysPage() {
async () => { async () => {
setIsDeleting(true); setIsDeleting(true);
try { try {
const response = await fetch('/connect/api-keys', { const response = await fetch("/connect/api-keys", {
method: 'DELETE', method: "DELETE",
headers: { headers: {
'Content-Type': 'application/json', "Content-Type": "application/json",
}, },
credentials: 'include', credentials: "include",
body: JSON.stringify({ id }), body: JSON.stringify({ id }),
}); });
if (!response.ok) { if (!response.ok) {
throw new Error('Failed to delete API key'); throw new Error("Failed to delete API key");
} }
setApiKeys(apiKeys.filter(key => key.id !== id)); setApiKeys(apiKeys.filter((key) => key.id !== id));
toast({ toast({
title: __("Success"), title: __("Success"),
description: __("API Key deleted successfully"), description: __("API Key deleted successfully"),
@@ -317,7 +336,9 @@ export default function APIKeysPage() {
} }
}, },
{ {
message: __(`Are you sure you want to delete the API key "${name}"? This action cannot be undone.`), message: __(
`Are you sure you want to delete the API key "${name}"? This action cannot be undone.`
),
} }
); );
}; };
@@ -326,11 +347,11 @@ export default function APIKeysPage() {
setIsLoadingKey(true); setIsLoadingKey(true);
try { try {
const response = await fetch(`/connect/api-keys/${id}`, { const response = await fetch(`/connect/api-keys/${id}`, {
credentials: 'include', credentials: "include",
}); });
if (!response.ok) { if (!response.ok) {
throw new Error('Failed to load API key'); throw new Error("Failed to load API key");
} }
const data: { key: string } = await response.json(); const data: { key: string } = await response.json();
@@ -400,7 +421,7 @@ export default function APIKeysPage() {
return ( return (
<Card key={apiKey.id} padded className="w-full"> <Card key={apiKey.id} padded className="w-full">
<div className="flex items-start justify-between gap-4"> <div className="flex items-start justify-between gap-4">
<div className="flex-1 space-y-2"> <div className="flex-1 space-y-2">
<div className="flex items-center gap-2"> <div className="flex items-center gap-2">
<h3 className="text-lg font-semibold">{apiKey.name}</h3> <h3 className="text-lg font-semibold">{apiKey.name}</h3>
@@ -464,7 +485,9 @@ export default function APIKeysPage() {
{__("Create an API key")} {__("Create an API key")}
</h2> </h2>
<p className="text-txt-tertiary mb-4"> <p className="text-txt-tertiary mb-4">
{__("Generate a new API key for programmatic access to your organization")} {__(
"Generate a new API key for programmatic access to your organization"
)}
</p> </p>
<Button <Button
onClick={() => dialogRef.current?.open()} onClick={() => dialogRef.current?.open()}
@@ -516,31 +539,34 @@ export default function APIKeysPage() {
type="button" type="button"
variant="tertiary" variant="tertiary"
onClick={() => { onClick={() => {
const allSelected = selectedOrganizations.length === organizations.length; const allSelected =
selectedOrganizations.length === organizations.length;
if (allSelected) { if (allSelected) {
setSelectedOrganizations([]); setSelectedOrganizations([]);
setOrganizationRoles({}); setOrganizationRoles({});
setValue("organizations", []); setValue("organizations", []);
} else { } else {
const allOrgIds = organizations.map(org => org.id); const allOrgIds = organizations.map((org) => org.id);
const newRoles: Record<string, string> = {}; const newRoles: Record<string, string> = {};
allOrgIds.forEach(id => { allOrgIds.forEach((id) => {
newRoles[id] = organizationRoles[id] || "FULL"; newRoles[id] = organizationRoles[id] || "FULL";
}); });
setSelectedOrganizations(allOrgIds); setSelectedOrganizations(allOrgIds);
setOrganizationRoles(newRoles); setOrganizationRoles(newRoles);
setValue( setValue(
"organizations", "organizations",
allOrgIds.map(id => ({ allOrgIds.map((id) => ({
organizationId: id, organizationId: id,
role: newRoles[id] role: newRoles[id],
})) }))
); );
} }
}} }}
className="text-xs h-7 min-h-7" className="text-xs h-7 min-h-7"
> >
{selectedOrganizations.length === organizations.length ? __("Clear All") : __("Select All")} {selectedOrganizations.length === organizations.length
? __("Clear All")
: __("Select All")}
</Button> </Button>
)} )}
</div> </div>
@@ -556,9 +582,7 @@ export default function APIKeysPage() {
<Th>{__("Name")}</Th> <Th>{__("Name")}</Th>
<Th width={180}>{__("Role")}</Th> <Th width={180}>{__("Role")}</Th>
<Th width={100}> <Th width={100}>
<div className="flex justify-end"> <div className="flex justify-end">{__("Access")}</div>
{__("Access")}
</div>
</Th> </Th>
</Tr> </Tr>
</Thead> </Thead>
@@ -570,30 +594,33 @@ export default function APIKeysPage() {
{org.name} {org.name}
</div> </div>
</Td> </Td>
<Td> <Td>
<div className="min-h-[36px] flex items-center"> <div className="min-h-[36px] flex items-center">
{selectedOrganizations.includes(org.id) ? ( {selectedOrganizations.includes(org.id) ? (
<Select <Select
value={organizationRoles[org.id] || "FULL"} value={organizationRoles[org.id] || "FULL"}
onValueChange={(role) => { onValueChange={(role) => {
const newRoles = { ...organizationRoles, [org.id]: role }; const newRoles = {
setOrganizationRoles(newRoles); ...organizationRoles,
setValue( [org.id]: role,
"organizations", };
selectedOrganizations.map(id => ({ setOrganizationRoles(newRoles);
organizationId: id, setValue(
role: newRoles[id] || "FULL" "organizations",
})) selectedOrganizations.map((id) => ({
); organizationId: id,
}} role: newRoles[id] || "FULL",
> }))
<Option value="FULL">{__("Full")}</Option> );
</Select> }}
) : ( >
<span className="text-txt-tertiary">—</span> <Option value="FULL">{__("Full")}</Option>
)} </Select>
</div> ) : (
</Td> <span className="text-txt-tertiary">—</span>
)}
</div>
</Td>
<Td> <Td>
<div className="flex justify-end"> <div className="flex justify-end">
<Checkbox <Checkbox
@@ -602,21 +629,26 @@ export default function APIKeysPage() {
let newSelected: string[]; let newSelected: string[];
const newRoles = { ...organizationRoles }; const newRoles = { ...organizationRoles };
if (checked) { if (checked) {
newSelected = [...selectedOrganizations, org.id]; newSelected = [
...selectedOrganizations,
org.id,
];
if (!newRoles[org.id]) { if (!newRoles[org.id]) {
newRoles[org.id] = "FULL"; newRoles[org.id] = "FULL";
} }
} else { } else {
newSelected = selectedOrganizations.filter(id => id !== org.id); newSelected = selectedOrganizations.filter(
(id) => id !== org.id
);
delete newRoles[org.id]; delete newRoles[org.id];
} }
setSelectedOrganizations(newSelected); setSelectedOrganizations(newSelected);
setOrganizationRoles(newRoles); setOrganizationRoles(newRoles);
setValue( setValue(
"organizations", "organizations",
newSelected.map(id => ({ newSelected.map((id) => ({
organizationId: id, organizationId: id,
role: newRoles[id] || "FULL" role: newRoles[id] || "FULL",
})) }))
); );
}} }}
@@ -660,14 +692,15 @@ export default function APIKeysPage() {
type="button" type="button"
variant="tertiary" variant="tertiary"
onClick={() => { onClick={() => {
const allSelected = selectedOrganizations.length === organizations.length; const allSelected =
selectedOrganizations.length === organizations.length;
if (allSelected) { if (allSelected) {
setSelectedOrganizations([]); setSelectedOrganizations([]);
setOrganizationRoles({}); setOrganizationRoles({});
} else { } else {
const allOrgIds = organizations.map(org => org.id); const allOrgIds = organizations.map((org) => org.id);
const newRoles: Record<string, string> = {}; const newRoles: Record<string, string> = {};
allOrgIds.forEach(id => { allOrgIds.forEach((id) => {
newRoles[id] = organizationRoles[id] || "FULL"; newRoles[id] = organizationRoles[id] || "FULL";
}); });
setSelectedOrganizations(allOrgIds); setSelectedOrganizations(allOrgIds);
@@ -676,7 +709,9 @@ export default function APIKeysPage() {
}} }}
className="text-xs h-7 min-h-7" className="text-xs h-7 min-h-7"
> >
{selectedOrganizations.length === organizations.length ? __("Clear All") : __("Select All")} {selectedOrganizations.length === organizations.length
? __("Clear All")
: __("Select All")}
</Button> </Button>
)} )}
</div> </div>
@@ -692,9 +727,7 @@ export default function APIKeysPage() {
<Th>{__("Name")}</Th> <Th>{__("Name")}</Th>
<Th width={180}>{__("Role")}</Th> <Th width={180}>{__("Role")}</Th>
<Th width={100}> <Th width={100}>
<div className="flex justify-end"> <div className="flex justify-end">{__("Access")}</div>
{__("Access")}
</div>
</Th> </Th>
</Tr> </Tr>
</Thead> </Thead>
@@ -712,7 +745,10 @@ export default function APIKeysPage() {
<Select <Select
value={organizationRoles[org.id] || "FULL"} value={organizationRoles[org.id] || "FULL"}
onValueChange={(role) => { onValueChange={(role) => {
const newRoles = { ...organizationRoles, [org.id]: role }; const newRoles = {
...organizationRoles,
[org.id]: role,
};
setOrganizationRoles(newRoles); setOrganizationRoles(newRoles);
}} }}
> >
@@ -731,12 +767,17 @@ export default function APIKeysPage() {
let newSelected: string[]; let newSelected: string[];
const newRoles = { ...organizationRoles }; const newRoles = { ...organizationRoles };
if (checked) { if (checked) {
newSelected = [...selectedOrganizations, org.id]; newSelected = [
...selectedOrganizations,
org.id,
];
if (!newRoles[org.id]) { if (!newRoles[org.id]) {
newRoles[org.id] = "FULL"; newRoles[org.id] = "FULL";
} }
} else { } else {
newSelected = selectedOrganizations.filter(id => id !== org.id); newSelected = selectedOrganizations.filter(
(id) => id !== org.id
);
delete newRoles[org.id]; delete newRoles[org.id];
} }
setSelectedOrganizations(newSelected); setSelectedOrganizations(newSelected);
@@ -754,7 +795,13 @@ export default function APIKeysPage() {
</div> </div>
</DialogContent> </DialogContent>
<DialogFooter> <DialogFooter>
<Button onClick={handleUpdate} disabled={isUpdating || (selectedOrganizations.length === 0 && !editingName.trim())}> <Button
onClick={handleUpdate}
disabled={
isUpdating ||
(selectedOrganizations.length === 0 && !editingName.trim())
}
>
{isUpdating ? __("Updating...") : __("Update")} {isUpdating ? __("Updating...") : __("Update")}
</Button> </Button>
</DialogFooter> </DialogFooter>

View File

@@ -10,7 +10,6 @@ import {
IconCircleCheck, IconCircleCheck,
} from "@probo/ui"; } from "@probo/ui";
import { PDFPreview } from "../components/documents/PDFPreview"; import { PDFPreview } from "../components/documents/PDFPreview";
import { buildEndpoint } from "/providers/RelayProviders";
import { useWindowSize } from "usehooks-ts"; import { useWindowSize } from "usehooks-ts";
import clsx from "clsx"; import clsx from "clsx";
import { sprintf } from "@probo/helpers"; import { sprintf } from "@probo/helpers";
@@ -63,7 +62,7 @@ export default function DocumentSigningRequestsPage() {
async function fetchDocuments() { async function fetchDocuments() {
try { try {
const response = await fetch( const response = await fetch(
buildEndpoint("/api/console/v1/documents/signing-requests"), "/api/console/v1/documents/signing-requests",
{ {
method: "GET", method: "GET",
headers: { headers: {
@@ -114,9 +113,7 @@ export default function DocumentSigningRequestsPage() {
setSigning(true); setSigning(true);
try { try {
const response = await fetch( const response = await fetch(
buildEndpoint( `/api/console/v1/documents/signing-requests/${docToSign.document_version_id}/sign`,
`/api/console/v1/documents/signing-requests/${docToSign.document_version_id}/sign`
),
{ {
method: "POST", method: "POST",
headers: { headers: {
@@ -235,9 +232,7 @@ export default function DocumentSigningRequestsPage() {
// Build PDF URL with watermark // Build PDF URL with watermark
const pdfUrl = token const pdfUrl = token
? `${buildEndpoint( ? `${`/api/console/v1/documents/signing-requests/${currentDoc.document_version_id}/pdf`}?token=${encodeURIComponent(token)}`
`/api/console/v1/documents/signing-requests/${currentDoc.document_version_id}/pdf`
)}?token=${encodeURIComponent(token)}`
: null; : null;
return ( return (
@@ -250,7 +245,10 @@ export default function DocumentSigningRequestsPage() {
<div className="grid lg:grid-cols-2 min-h-0 h-full"> <div className="grid lg:grid-cols-2 min-h-0 h-full">
<div className="max-w-[440px] mx-auto py-20 overflow-y-auto scrollbar-hide"> <div className="max-w-[440px] mx-auto py-20 overflow-y-auto scrollbar-hide">
<h1 className="text-2xl font-semibold mb-6"> <h1 className="text-2xl font-semibold mb-6">
{sprintf(__("%s requests your signature"), signingData.organizationName)} {sprintf(
__("%s requests your signature"),
signingData.organizationName
)}
</h1> </h1>
{allSigned ? ( {allSigned ? (
<p className="text-txt-secondary text-base"> <p className="text-txt-secondary text-base">
@@ -266,7 +264,10 @@ export default function DocumentSigningRequestsPage() {
<Card className="mb-6 overflow-hidden"> <Card className="mb-6 overflow-hidden">
<div className="divide-y divide-border-solid"> <div className="divide-y divide-border-solid">
{(() => { {(() => {
const renderDocumentItem = (doc: Document, index: number) => ( const renderDocumentItem = (
doc: Document,
index: number
) => (
<div <div
key={doc.document_version_id} key={doc.document_version_id}
className={clsx( className={clsx(
@@ -307,15 +308,15 @@ export default function DocumentSigningRequestsPage() {
doc.signed doc.signed
? "bg-green-100 text-green-800" ? "bg-green-100 text-green-800"
: index === currentDocIndex : index === currentDocIndex
? "bg-blue-100 text-blue-800" ? "bg-blue-100 text-blue-800"
: "bg-gray-100 text-gray-700" : "bg-gray-100 text-gray-700"
)} )}
> >
{doc.signed {doc.signed
? __("Signed") ? __("Signed")
: index === currentDocIndex : index === currentDocIndex
? __("In review") ? __("In review")
: __("Waiting signature")} : __("Waiting signature")}
</span> </span>
</div> </div>
</div> </div>
@@ -352,8 +353,12 @@ export default function DocumentSigningRequestsPage() {
const currentIsLast = currentDocIndex === totalDocs - 1; const currentIsLast = currentDocIndex === totalDocs - 1;
// Calculate hidden docs before and after current // Calculate hidden docs before and after current
const hiddenBeforeCurrent = currentIsFirst ? 0 : currentDocIndex - 1; const hiddenBeforeCurrent = currentIsFirst
const hiddenAfterCurrent = currentIsLast ? 0 : totalDocs - currentDocIndex - 2; ? 0
: currentDocIndex - 1;
const hiddenAfterCurrent = currentIsLast
? 0
: totalDocs - currentDocIndex - 2;
return ( return (
<> <>
@@ -367,12 +372,16 @@ export default function DocumentSigningRequestsPage() {
className="w-full py-3 px-4 text-sm text-txt-secondary hover:bg-level-1 transition-colors text-left flex items-center gap-2" className="w-full py-3 px-4 text-sm text-txt-secondary hover:bg-level-1 transition-colors text-left flex items-center gap-2"
> >
<span className="text-txt-tertiary">•••</span> <span className="text-txt-tertiary">•••</span>
{sprintf(__("Show %s more documents"), hiddenBeforeCurrent)} {sprintf(
__("Show %s more documents"),
hiddenBeforeCurrent
)}
</button> </button>
)} )}
{/* Current document (if not first) */} {/* Current document (if not first) */}
{!currentIsFirst && renderDocumentItem(currentDoc, currentDocIndex)} {!currentIsFirst &&
renderDocumentItem(currentDoc, currentDocIndex)}
{/* Show more button for documents AFTER current (upcoming documents) */} {/* Show more button for documents AFTER current (upcoming documents) */}
{hiddenAfterCurrent > 0 && ( {hiddenAfterCurrent > 0 && (
@@ -381,7 +390,10 @@ export default function DocumentSigningRequestsPage() {
className="w-full py-3 px-4 text-sm text-txt-secondary hover:bg-level-1 transition-colors text-left flex items-center gap-2" className="w-full py-3 px-4 text-sm text-txt-secondary hover:bg-level-1 transition-colors text-left flex items-center gap-2"
> >
<span className="text-txt-tertiary">•••</span> <span className="text-txt-tertiary">•••</span>
{sprintf(__("Show %s more documents"), hiddenAfterCurrent)} {sprintf(
__("Show %s more documents"),
hiddenAfterCurrent
)}
</button> </button>
)} )}
</> </>
@@ -390,9 +402,7 @@ export default function DocumentSigningRequestsPage() {
</div> </div>
</Card> </Card>
<p className="text-txt-secondary text-sm mb-6"> <p className="text-txt-secondary text-sm mb-6">
{__( {__("Please review the document carefully before signing.")}
"Please review the document carefully before signing."
)}
</p> </p>
</> </>
)} )}
@@ -421,10 +431,7 @@ export default function DocumentSigningRequestsPage() {
</> </>
)} )}
{currentDoc.signed && !isLastDocument && ( {currentDoc.signed && !isLastDocument && (
<Button <Button onClick={handleNextDocument} className="h-10 w-full mt-4">
onClick={handleNextDocument}
className="h-10 w-full mt-4"
>
{__("Next Document")} {__("Next Document")}
</Button> </Button>
)} )}

View File

@@ -4,7 +4,6 @@ import { useTranslate } from "@probo/i18n";
import { z } from "zod"; import { z } from "zod";
import { useFormWithSchema } from "/hooks/useFormWithSchema"; import { useFormWithSchema } from "/hooks/useFormWithSchema";
import { usePageTitle } from "@probo/hooks"; import { usePageTitle } from "@probo/hooks";
import { buildEndpoint } from "/providers/RelayProviders";
import { useState } from "react"; import { useState } from "react";
const schema = z.object({ const schema = z.object({
@@ -22,17 +21,14 @@ export default function ForgotPasswordPage() {
}); });
const onSubmit = handleSubmit(async (data) => { const onSubmit = handleSubmit(async (data) => {
const response = await fetch( const response = await fetch("/connect/forget-password", {
buildEndpoint("/connect/forget-password"), method: "POST",
{ headers: {
method: "POST", "Content-Type": "application/json",
headers: { },
"Content-Type": "application/json", credentials: "include",
}, body: JSON.stringify(data),
credentials: "include", });
body: JSON.stringify(data),
}
);
if (!response.ok) { if (!response.ok) {
const errorData = await response.json().catch(() => ({})); const errorData = await response.json().catch(() => ({}));

View File

@@ -3,7 +3,6 @@ import { Button, Field, IconChevronLeft, useToast } from "@probo/ui";
import type { FormEventHandler } from "react"; import type { FormEventHandler } from "react";
import { useState } from "react"; import { useState } from "react";
import { Link, useSearchParams } from "react-router"; import { Link, useSearchParams } from "react-router";
import { buildEndpoint } from "/providers/RelayProviders";
export default function LoginPage() { export default function LoginPage() {
const { __ } = useTranslate(); const { __ } = useTranslate();
@@ -33,7 +32,7 @@ export default function LoginPage() {
setIsLoading(true); setIsLoading(true);
try { try {
const res = await fetch(buildEndpoint("/connect/login"), { const res = await fetch("/connect/login", {
method: "POST", method: "POST",
headers: { headers: {
"Content-Type": "application/json", "Content-Type": "application/json",
@@ -68,7 +67,7 @@ export default function LoginPage() {
setIsChecking(true); setIsChecking(true);
try { try {
const res = await fetch(buildEndpoint("/connect/check-sso"), { const res = await fetch("/connect/check-sso", {
method: "POST", method: "POST",
headers: { headers: {
"Content-Type": "application/json", "Content-Type": "application/json",
@@ -78,15 +77,15 @@ export default function LoginPage() {
if (!res.ok) { if (!res.ok) {
const error = await res.json(); const error = await res.json();
throw new Error(error.message || __("SSO not available for this email domain")); throw new Error(
error.message || __("SSO not available for this email domain")
);
} }
const data = await res.json(); const data = await res.json();
if (data.ssoAvailable && data.samlConfigId) { if (data.ssoAvailable && data.samlConfigId) {
window.location.href = buildEndpoint( window.location.href = `/connect/saml/login/${data.samlConfigId}`;
`/connect/saml/login/${data.samlConfigId}`
);
} else { } else {
throw new Error(__("SSO not available for this email domain")); throw new Error(__("SSO not available for this email domain"));
} }
@@ -115,10 +114,7 @@ export default function LoginPage() {
{__("Choose your login method")} {__("Choose your login method")}
</p> </p>
<Button <Button className="w-full" onClick={() => setMode("password")}>
className="w-full"
onClick={() => setMode("password")}
>
{__("Login with Email")} {__("Login with Email")}
</Button> </Button>
@@ -146,7 +142,10 @@ export default function LoginPage() {
<div className="text-center mt-6 text-sm text-txt-secondary"> <div className="text-center mt-6 text-sm text-txt-secondary">
{__("Don't have an account ?")}{" "} {__("Don't have an account ?")}{" "}
<Link to="/auth/register" className="underline hover:text-txt-primary"> <Link
to="/auth/register"
className="underline hover:text-txt-primary"
>
{__("Register")} {__("Register")}
</Link> </Link>
</div> </div>
@@ -206,7 +205,10 @@ export default function LoginPage() {
<div className="text-center mt-6 text-sm text-txt-secondary"> <div className="text-center mt-6 text-sm text-txt-secondary">
{__("Don't have an account ?")}{" "} {__("Don't have an account ?")}{" "}
<Link to="/auth/register" className="underline hover:text-txt-primary"> <Link
to="/auth/register"
className="underline hover:text-txt-primary"
>
{__("Register")} {__("Register")}
</Link> </Link>
</div> </div>
@@ -235,9 +237,7 @@ export default function LoginPage() {
<span className="text-sm">{__("Back")}</span> <span className="text-sm">{__("Back")}</span>
</button> </button>
<h1 className="text-center text-2xl font-bold"> <h1 className="text-center text-2xl font-bold">{__("Login with SSO")}</h1>
{__("Login with SSO")}
</h1>
<p className="text-center text-txt-tertiary mt-1 mb-6"> <p className="text-center text-txt-tertiary mt-1 mb-6">
{__("Enter your work email to continue with SSO")} {__("Enter your work email to continue with SSO")}
</p> </p>

View File

@@ -4,7 +4,6 @@ import { useTranslate } from "@probo/i18n";
import { z } from "zod"; import { z } from "zod";
import { useFormWithSchema } from "/hooks/useFormWithSchema"; import { useFormWithSchema } from "/hooks/useFormWithSchema";
import { usePageTitle } from "@probo/hooks"; import { usePageTitle } from "@probo/hooks";
import { buildEndpoint } from "/providers/RelayProviders";
const schema = z.object({ const schema = z.object({
email: z.string().email(), email: z.string().email(),
@@ -25,17 +24,14 @@ export default function RegisterPage() {
}); });
const onSubmit = handleSubmit(async (data) => { const onSubmit = handleSubmit(async (data) => {
const response = await fetch( const response = await fetch("/connect/register", {
buildEndpoint("/connect/register"), method: "POST",
{ headers: {
method: "POST", "Content-Type": "application/json",
headers: {
"Content-Type": "application/json",
},
credentials: "include",
body: JSON.stringify(data),
}, },
); credentials: "include",
body: JSON.stringify(data),
});
// Registration failed // Registration failed
if (!response.ok) { if (!response.ok) {

View File

@@ -4,7 +4,6 @@ import { useTranslate } from "@probo/i18n";
import { z } from "zod"; import { z } from "zod";
import { useFormWithSchema } from "/hooks/useFormWithSchema"; import { useFormWithSchema } from "/hooks/useFormWithSchema";
import { usePageTitle } from "@probo/hooks"; import { usePageTitle } from "@probo/hooks";
import { buildEndpoint } from "/providers/RelayProviders";
const schema = z const schema = z
.object({ .object({
@@ -20,15 +19,12 @@ export default function ResetPasswordPage() {
const { __ } = useTranslate(); const { __ } = useTranslate();
const navigate = useNavigate(); const navigate = useNavigate();
const { toast } = useToast(); const { toast } = useToast();
const { register, handleSubmit, formState } = useFormWithSchema( const { register, handleSubmit, formState } = useFormWithSchema(schema, {
schema, defaultValues: {
{ password: "",
defaultValues: { confirmPassword: "",
password: "", },
confirmPassword: "", });
},
}
);
const onSubmit = handleSubmit(async (data) => { const onSubmit = handleSubmit(async (data) => {
const searchParams = new URLSearchParams(location.search); const searchParams = new URLSearchParams(location.search);
@@ -43,20 +39,17 @@ export default function ResetPasswordPage() {
return; return;
} }
const response = await fetch( const response = await fetch("/connect/reset-password", {
buildEndpoint("/connect/reset-password"), method: "POST",
{ headers: {
method: "POST", "Content-Type": "application/json",
headers: { },
"Content-Type": "application/json", credentials: "include",
}, body: JSON.stringify({
credentials: "include", token: token,
body: JSON.stringify({ password: data.password,
token: token, }),
password: data.password, });
}),
}
);
// Reset failed // Reset failed
if (!response.ok) { if (!response.ok) {

View File

@@ -4,7 +4,6 @@ import { useTranslate } from "@probo/i18n";
import { z } from "zod"; import { z } from "zod";
import { useFormWithSchema } from "/hooks/useFormWithSchema"; import { useFormWithSchema } from "/hooks/useFormWithSchema";
import { usePageTitle } from "@probo/hooks"; import { usePageTitle } from "@probo/hooks";
import { buildEndpoint } from "/providers/RelayProviders";
import { useEffect } from "react"; import { useEffect } from "react";
const schema = z.object({ const schema = z.object({
@@ -50,21 +49,18 @@ export default function SignupFromInvitationPage() {
return; return;
} }
const response = await fetch( const response = await fetch("/connect/signup-from-invitation", {
buildEndpoint("/connect/signup-from-invitation"), method: "POST",
{ headers: {
method: "POST", "Content-Type": "application/json",
headers: { },
"Content-Type": "application/json", credentials: "include",
}, body: JSON.stringify({
credentials: "include", token: token,
body: JSON.stringify({ password: data.password,
token: token, fullName: data.fullName,
password: data.password, }),
fullName: data.fullName, });
}),
}
);
if (!response.ok) { if (!response.ok) {
const errorData = await response.json().catch(() => ({})); const errorData = await response.json().catch(() => ({}));
@@ -78,7 +74,9 @@ export default function SignupFromInvitationPage() {
toast({ toast({
title: __("Success"), title: __("Success"),
description: __("Account created successfully. Please accept your invitation to join the organization."), description: __(
"Account created successfully. Please accept your invitation to join the organization."
),
variant: "success", variant: "success",
}); });
navigate("/", { replace: true }); navigate("/", { replace: true });

View File

@@ -1,242 +0,0 @@
import {
Environment,
type FetchFunction,
Network,
RecordSource,
Store,
} from "relay-runtime";
import { GraphQLError } from "graphql";
import type { PropsWithChildren } from "react";
import { RelayEnvironmentProvider } from "react-relay";
export class UnAuthenticatedError extends Error {
constructor(message?: string) {
super(message || "UNAUTHENTICATED");
this.name = "UnAuthenticatedError";
}
}
export class InternalServerError extends Error {
constructor() {
super("INTERNAL_SERVER_ERROR");
this.name = "InternalServerError";
}
}
export class AuthenticationRequiredError extends Error {
public redirectUrl: string;
public requiresSaml: boolean;
public organizationId: string;
public samlConfigId?: string;
constructor(extensions: {
redirectUrl: string;
requiresSaml: boolean;
organizationId: string;
samlConfigId?: string;
}) {
super("AUTHENTICATION_REQUIRED");
this.name = "AuthenticationRequiredError";
this.redirectUrl = extensions.redirectUrl;
this.requiresSaml = extensions.requiresSaml;
this.organizationId = extensions.organizationId;
this.samlConfigId = extensions.samlConfigId;
}
}
export class UnauthorizedError extends Error {
constructor(message?: string) {
super(message || "UNAUTHORIZED");
this.name = "UnauthorizedError";
}
}
export class ForbiddenError extends Error {
constructor(message?: string) {
super(message || "FORBIDDEN");
this.name = "ForbiddenError";
}
}
export class InvalidError extends Error {
field?: string;
cause?: string;
constructor(message?: string, field?: string, cause?: string) {
super(message || "INVALID");
this.name = "InvalidError";
this.field = field;
this.cause = cause;
}
}
export function buildEndpoint(path: string): string {
const host = import.meta.env.VITE_API_URL;
if (!host) {
return path;
}
const formattedHost =
host.startsWith("http://") || host.startsWith("https://")
? host
: `https://${host}`;
const url = new URL(formattedHost);
if (path) {
url.pathname = path.startsWith("/") ? path : `/${path}`;
}
return url.toString();
}
const hasUnauthenticatedError = (error: GraphQLError) =>
error.extensions?.code == "UNAUTHENTICATED";
const hasAuthenticationRequiredError = (error: GraphQLError) =>
error.extensions?.code == "AUTHENTICATION_REQUIRED";
const hasUnauthorizedError = (error: GraphQLError) =>
error.extensions?.code == "UNAUTHORIZED";
const hasForbiddenError = (error: GraphQLError) =>
error.extensions?.code == "FORBIDDEN";
const hasInvalidError = (error: GraphQLError) =>
error.extensions?.code == "INVALID_REQUEST";
const fetchRelay: FetchFunction = async (
request,
variables,
_,
uploadables
) => {
const requestInit: RequestInit = {
method: "POST",
credentials: "include",
headers: {},
};
if (uploadables) {
const formData = new FormData();
formData.append(
"operations",
JSON.stringify({
operationName: request.name,
query: request.text,
variables: variables,
})
);
const uploadableMap: {
[key: string]: string[];
} = {};
Object.keys(uploadables).forEach((key, index) => {
uploadableMap[index] = [`variables.${key}`];
});
formData.append("map", JSON.stringify(uploadableMap));
Object.keys(uploadables).forEach((key, index) => {
formData.append(index.toString(), uploadables[key]);
});
requestInit.body = formData;
} else {
requestInit.headers = {
Accept:
"application/graphql-response+json; charset=utf-8, application/json; charset=utf-8",
"Content-Type": "application/json",
};
requestInit.body = JSON.stringify({
operationName: request.name,
query: request.text,
variables,
});
}
const response = await fetch(
buildEndpoint("/api/console/v1/query"),
requestInit
);
if (response.status === 500) {
throw new InternalServerError();
}
const json = await response.json();
if (json.errors) {
const errors = json.errors as GraphQLError[];
const unauthenticatedError = errors.find(hasUnauthenticatedError);
if (unauthenticatedError) {
throw new UnAuthenticatedError(unauthenticatedError.message);
}
const authRequiredError = errors.find(hasAuthenticationRequiredError);
if (authRequiredError?.extensions) {
const { redirectUrl, requiresSaml, organizationId, samlConfigId } = authRequiredError.extensions;
throw new AuthenticationRequiredError({
redirectUrl: redirectUrl as string,
requiresSaml: requiresSaml as boolean,
organizationId: organizationId as string,
samlConfigId: samlConfigId as string | undefined,
});
}
const unauthorizedError = errors.find(hasUnauthorizedError);
if (unauthorizedError) {
throw new UnauthorizedError(unauthorizedError.message);
}
const forbiddenError = errors.find(hasForbiddenError);
if (forbiddenError) {
throw new ForbiddenError(forbiddenError.message);
}
const invalidError = errors.find(hasInvalidError);
if (invalidError) {
throw new InvalidError(
invalidError.message,
invalidError.extensions.field as string ?? "",
invalidError.extensions.cause as string ?? "",
);
}
}
return json;
};
const source = new RecordSource();
const store = new Store(source, {
queryCacheExpirationTime: 1 * 60 * 1000,
gcReleaseBufferSize: 20,
});
export const relayEnvironment = new Environment({
network: Network.create(fetchRelay),
store,
});
export const clearRelayStore = () => {
const source = relayEnvironment.getStore().getSource();
if (source instanceof Map) {
source.clear();
}
};
/**
* Provider for relay with the probo environment
*/
export function RelayProvider({ children }: PropsWithChildren) {
return (
<RelayEnvironmentProvider environment={relayEnvironment}>
{children}
</RelayEnvironmentProvider>
);
}

View File

@@ -7,12 +7,6 @@ import {
import { MainLayout } from "./layouts/MainLayout"; import { MainLayout } from "./layouts/MainLayout";
import { EmployeeLayout } from "./layouts/EmployeeLayout"; import { EmployeeLayout } from "./layouts/EmployeeLayout";
import { AuthLayout, CenteredLayout, CenteredLayoutSkeleton } from "@probo/ui"; import { AuthLayout, CenteredLayout, CenteredLayoutSkeleton } from "@probo/ui";
import {
relayEnvironment,
UnAuthenticatedError,
UnauthorizedError,
ForbiddenError,
} from "./providers/RelayProviders";
import { PageSkeleton } from "./components/skeletons/PageSkeleton.tsx"; import { PageSkeleton } from "./components/skeletons/PageSkeleton.tsx";
import { loadQuery } from "react-relay"; import { loadQuery } from "react-relay";
import { riskRoutes } from "./routes/riskRoutes.ts"; import { riskRoutes } from "./routes/riskRoutes.ts";
@@ -37,12 +31,24 @@ import { rightsRequestRoutes } from "./routes/rightsRequestRoutes.ts";
import { processingActivityRoutes } from "./routes/processingActivityRoutes.ts"; import { processingActivityRoutes } from "./routes/processingActivityRoutes.ts";
import { statesOfApplicabilityRoutes } from "./routes/statesOfApplicabilityRoutes.ts"; import { statesOfApplicabilityRoutes } from "./routes/statesOfApplicabilityRoutes.ts";
import { lazy } from "@probo/react-lazy"; import { lazy } from "@probo/react-lazy";
import { loaderFromQueryLoader, routeFromAppRoute, withQueryRef, type AppRoute } from "@probo/routes"; import {
loaderFromQueryLoader,
routeFromAppRoute,
withQueryRef,
type AppRoute,
} from "@probo/routes";
import { employeeDocumentsQuery } from "./pages/organizations/employee/EmployeeDocumentsPage"; import { employeeDocumentsQuery } from "./pages/organizations/employee/EmployeeDocumentsPage";
import { employeeDocumentSignatureQuery } from "./pages/organizations/employee/EmployeeDocumentSignaturePage"; import { employeeDocumentSignatureQuery } from "./pages/organizations/employee/EmployeeDocumentSignaturePage";
import { Role } from "@probo/helpers"; import { Role } from "@probo/helpers";
import { PermissionsContext } from "./providers/PermissionsContext"; import { PermissionsContext } from "./providers/PermissionsContext";
import { use } from "react"; import { use } from "react";
import { RelayEnvironmentProvider } from "react-relay";
import { connectEnvironment, consoleEnvironment } from "./environments.ts";
import {
ForbiddenError,
UnAuthenticatedError,
UnauthorizedError,
} from "@probo/relay";
/** /**
* Top level error boundary * Top level error boundary
@@ -68,7 +74,11 @@ function ErrorBoundary({ error: propsError }: { error?: string }) {
const routes = [ const routes = [
{ {
path: "/auth", path: "/auth",
Component: AuthLayout, Component: () => (
<RelayEnvironmentProvider environment={connectEnvironment}>
<AuthLayout />
</RelayEnvironmentProvider>
),
children: [ children: [
{ {
path: "login", path: "login",
@@ -98,7 +108,11 @@ const routes = [
}, },
{ {
path: "/", path: "/",
Component: CenteredLayout, Component: () => (
<RelayEnvironmentProvider environment={consoleEnvironment}>
<CenteredLayout />
</RelayEnvironmentProvider>
),
Fallback: CenteredLayoutSkeleton, Fallback: CenteredLayoutSkeleton,
ErrorBoundary: ErrorBoundary, ErrorBoundary: ErrorBoundary,
children: [ children: [
@@ -132,35 +146,42 @@ const routes = [
{ {
path: "", path: "",
Fallback: PageSkeleton, Fallback: PageSkeleton,
loader: loaderFromQueryLoader( loader: loaderFromQueryLoader(({ organizationId }) =>
({ organizationId }) => loadQuery(consoleEnvironment, employeeDocumentsQuery, {
loadQuery(relayEnvironment, employeeDocumentsQuery, { organizationId: organizationId!,
organizationId: organizationId!, })
}) ),
Component: withQueryRef(
lazy(
() => import("./pages/organizations/employee/EmployeeDocumentsPage")
)
), ),
Component: withQueryRef(lazy(
() => import("./pages/organizations/employee/EmployeeDocumentsPage")
)),
}, },
{ {
path: ":documentId", path: ":documentId",
Fallback: PageSkeleton, Fallback: PageSkeleton,
ErrorBoundary: ErrorBoundary, ErrorBoundary: ErrorBoundary,
loader: loaderFromQueryLoader( loader: loaderFromQueryLoader(({ documentId }) =>
({ documentId }) => loadQuery(consoleEnvironment, employeeDocumentSignatureQuery, {
loadQuery(relayEnvironment, employeeDocumentSignatureQuery, { documentId: documentId!,
documentId: documentId!, })
}) ),
Component: withQueryRef(
lazy(
() =>
import("./pages/organizations/employee/EmployeeDocumentSignaturePage")
)
), ),
Component: withQueryRef(lazy(
() => import("./pages/organizations/employee/EmployeeDocumentSignaturePage")
)),
}, },
], ],
}, },
{ {
path: "/organizations/:organizationId", path: "/organizations/:organizationId",
Component: MainLayout, Component: () => (
<RelayEnvironmentProvider environment={consoleEnvironment}>
<MainLayout />
</RelayEnvironmentProvider>
),
ErrorBoundary: ErrorBoundary, ErrorBoundary: ErrorBoundary,
children: [ children: [
{ {
@@ -180,13 +201,14 @@ const routes = [
{ {
path: "settings", path: "settings",
Fallback: PageSkeleton, Fallback: PageSkeleton,
loader: loaderFromQueryLoader( loader: loaderFromQueryLoader(({ organizationId }) =>
({ organizationId }) => loadQuery(consoleEnvironment, organizationViewQuery, {
loadQuery(relayEnvironment, organizationViewQuery, { organizationId: organizationId!,
organizationId: organizationId!, })
}) ),
Component: withQueryRef(
lazy(() => import("./pages/organizations/SettingsPage"))
), ),
Component: withQueryRef(lazy(() => import("./pages/organizations/SettingsPage"))),
children: [ children: [
{ {
path: "", path: "",
@@ -196,19 +218,27 @@ const routes = [
}, },
{ {
path: "general", path: "general",
Component: lazy(() => import("./pages/organizations/settings/GeneralSettingsTab")), Component: lazy(
() => import("./pages/organizations/settings/GeneralSettingsTab")
),
}, },
{ {
path: "members", path: "members",
Component: lazy(() => import("./pages/organizations/settings/MembersSettingsTab")), Component: lazy(
() => import("./pages/organizations/settings/MembersSettingsTab")
),
}, },
{ {
path: "domain", path: "domain",
Component: lazy(() => import("./pages/organizations/settings/DomainSettingsTab")), Component: lazy(
() => import("./pages/organizations/settings/DomainSettingsTab")
),
}, },
{ {
path: "saml-sso", path: "saml-sso",
Component: lazy(() => import("./pages/organizations/settings/SAMLSettingsTab")), Component: lazy(
() => import("./pages/organizations/settings/SAMLSettingsTab")
),
}, },
], ],
}, },

View File

@@ -1,5 +1,5 @@
import { loadQuery } from "react-relay"; import { loadQuery } from "react-relay";
import { relayEnvironment } from "/providers/RelayProviders"; import { consoleEnvironment } from "/environments";
import { PageSkeleton } from "/components/skeletons/PageSkeleton"; import { PageSkeleton } from "/components/skeletons/PageSkeleton";
import { lazy } from "@probo/react-lazy"; import { lazy } from "@probo/react-lazy";
import { assetsQuery, assetNodeQuery } from "../hooks/graph/AssetGraph"; import { assetsQuery, assetNodeQuery } from "../hooks/graph/AssetGraph";
@@ -12,7 +12,7 @@ export const assetRoutes = [
path: "assets", path: "assets",
Fallback: PageSkeleton, Fallback: PageSkeleton,
loader: loaderFromQueryLoader(({ organizationId }) => loader: loaderFromQueryLoader(({ organizationId }) =>
loadQuery<AssetGraphListQuery>(relayEnvironment, assetsQuery, { loadQuery<AssetGraphListQuery>(consoleEnvironment, assetsQuery, {
organizationId: organizationId, organizationId: organizationId,
snapshotId: null, snapshotId: null,
}), }),
@@ -23,7 +23,7 @@ export const assetRoutes = [
path: "snapshots/:snapshotId/assets", path: "snapshots/:snapshotId/assets",
Fallback: PageSkeleton, Fallback: PageSkeleton,
loader: loaderFromQueryLoader(({ organizationId, snapshotId }) => loader: loaderFromQueryLoader(({ organizationId, snapshotId }) =>
loadQuery<AssetGraphListQuery>(relayEnvironment, assetsQuery, { loadQuery<AssetGraphListQuery>(consoleEnvironment, assetsQuery, {
organizationId: organizationId, organizationId: organizationId,
snapshotId: snapshotId, snapshotId: snapshotId,
}), }),
@@ -34,7 +34,7 @@ export const assetRoutes = [
path: "assets/:assetId", path: "assets/:assetId",
Fallback: PageSkeleton, Fallback: PageSkeleton,
loader: loaderFromQueryLoader(({ assetId }) => loader: loaderFromQueryLoader(({ assetId }) =>
loadQuery<AssetGraphNodeQuery>(relayEnvironment, assetNodeQuery, { assetId }), loadQuery<AssetGraphNodeQuery>(consoleEnvironment, assetNodeQuery, { assetId }),
), ),
Component: withQueryRef(lazy( Component: withQueryRef(lazy(
() => import("/pages/organizations/assets/AssetDetailsPage"), () => import("/pages/organizations/assets/AssetDetailsPage"),
@@ -44,7 +44,7 @@ export const assetRoutes = [
path: "snapshots/:snapshotId/assets/:assetId", path: "snapshots/:snapshotId/assets/:assetId",
Fallback: PageSkeleton, Fallback: PageSkeleton,
loader: loaderFromQueryLoader(({ assetId }) => loader: loaderFromQueryLoader(({ assetId }) =>
loadQuery<AssetGraphNodeQuery>(relayEnvironment, assetNodeQuery, { assetId }), loadQuery<AssetGraphNodeQuery>(consoleEnvironment, assetNodeQuery, { assetId }),
), ),
Component: withQueryRef(lazy( Component: withQueryRef(lazy(
() => import("/pages/organizations/assets/AssetDetailsPage"), () => import("/pages/organizations/assets/AssetDetailsPage"),

View File

@@ -1,5 +1,5 @@
import { loadQuery } from "react-relay"; import { loadQuery } from "react-relay";
import { relayEnvironment } from "/providers/RelayProviders"; import { consoleEnvironment } from "/environments";
import { PageSkeleton } from "/components/skeletons/PageSkeleton"; import { PageSkeleton } from "/components/skeletons/PageSkeleton";
import { lazy } from "@probo/react-lazy"; import { lazy } from "@probo/react-lazy";
import { auditsQuery, auditNodeQuery } from "../hooks/graph/AuditGraph"; import { auditsQuery, auditNodeQuery } from "../hooks/graph/AuditGraph";
@@ -12,7 +12,7 @@ export const auditRoutes = [
path: "audits", path: "audits",
Fallback: PageSkeleton, Fallback: PageSkeleton,
loader: loaderFromQueryLoader(({ organizationId }) => loader: loaderFromQueryLoader(({ organizationId }) =>
loadQuery<AuditGraphListQuery>(relayEnvironment, auditsQuery, { organizationId }), loadQuery<AuditGraphListQuery>(consoleEnvironment, auditsQuery, { organizationId }),
), ),
Component: withQueryRef(lazy( Component: withQueryRef(lazy(
() => import("/pages/organizations/audits/AuditsPage") () => import("/pages/organizations/audits/AuditsPage")
@@ -22,7 +22,7 @@ export const auditRoutes = [
path: "audits/:auditId", path: "audits/:auditId",
Fallback: PageSkeleton, Fallback: PageSkeleton,
loader: loaderFromQueryLoader(({ auditId }) => loader: loaderFromQueryLoader(({ auditId }) =>
loadQuery<AuditGraphNodeQuery>(relayEnvironment, auditNodeQuery, { auditId }), loadQuery<AuditGraphNodeQuery>(consoleEnvironment, auditNodeQuery, { auditId }),
), ),
Component: withQueryRef(lazy( Component: withQueryRef(lazy(
() => import("/pages/organizations/audits/AuditDetailsPage") () => import("/pages/organizations/audits/AuditDetailsPage")

View File

@@ -1,5 +1,5 @@
import { loadQuery } from "react-relay"; import { loadQuery } from "react-relay";
import { relayEnvironment } from "/providers/RelayProviders"; import { consoleEnvironment } from "/environments";
import { PageSkeleton } from "/components/skeletons/PageSkeleton"; import { PageSkeleton } from "/components/skeletons/PageSkeleton";
import { lazy } from "@probo/react-lazy"; import { lazy } from "@probo/react-lazy";
import { continualImprovementsQuery, continualImprovementNodeQuery } from "/hooks/graph/ContinualImprovementGraph"; import { continualImprovementsQuery, continualImprovementNodeQuery } from "/hooks/graph/ContinualImprovementGraph";
@@ -12,7 +12,7 @@ export const continualImprovementRoutes = [
path: "continual-improvements", path: "continual-improvements",
Fallback: PageSkeleton, Fallback: PageSkeleton,
loader: loaderFromQueryLoader(({ organizationId }) => loader: loaderFromQueryLoader(({ organizationId }) =>
loadQuery<ContinualImprovementGraphListQuery>(relayEnvironment, continualImprovementsQuery, { loadQuery<ContinualImprovementGraphListQuery>(consoleEnvironment, continualImprovementsQuery, {
organizationId, organizationId,
snapshotId: null snapshotId: null
}), }),
@@ -25,7 +25,7 @@ export const continualImprovementRoutes = [
path: "snapshots/:snapshotId/continual-improvements", path: "snapshots/:snapshotId/continual-improvements",
Fallback: PageSkeleton, Fallback: PageSkeleton,
loader: loaderFromQueryLoader(({ organizationId, snapshotId }) => loader: loaderFromQueryLoader(({ organizationId, snapshotId }) =>
loadQuery<ContinualImprovementGraphListQuery>(relayEnvironment, continualImprovementsQuery, { loadQuery<ContinualImprovementGraphListQuery>(consoleEnvironment, continualImprovementsQuery, {
organizationId, organizationId,
snapshotId, snapshotId,
}), }),
@@ -38,7 +38,7 @@ export const continualImprovementRoutes = [
path: "continual-improvements/:improvementId", path: "continual-improvements/:improvementId",
Fallback: PageSkeleton, Fallback: PageSkeleton,
loader: loaderFromQueryLoader(({ improvementId }) => loader: loaderFromQueryLoader(({ improvementId }) =>
loadQuery<ContinualImprovementGraphNodeQuery>(relayEnvironment, continualImprovementNodeQuery, { loadQuery<ContinualImprovementGraphNodeQuery>(consoleEnvironment, continualImprovementNodeQuery, {
continualImprovementId: improvementId!, continualImprovementId: improvementId!,
}), }),
), ),
@@ -50,7 +50,7 @@ export const continualImprovementRoutes = [
path: "snapshots/:snapshotId/continual-improvements/:improvementId", path: "snapshots/:snapshotId/continual-improvements/:improvementId",
Fallback: PageSkeleton, Fallback: PageSkeleton,
loader: loaderFromQueryLoader(({ improvementId }) => loader: loaderFromQueryLoader(({ improvementId }) =>
loadQuery<ContinualImprovementGraphNodeQuery>(relayEnvironment, continualImprovementNodeQuery, { loadQuery<ContinualImprovementGraphNodeQuery>(consoleEnvironment, continualImprovementNodeQuery, {
continualImprovementId: improvementId!, continualImprovementId: improvementId!,
}), }),
), ),

View File

@@ -1,5 +1,5 @@
import { loadQuery } from "react-relay"; import { loadQuery } from "react-relay";
import { relayEnvironment } from "/providers/RelayProviders"; import { consoleEnvironment } from "/environments";
import { PageSkeleton } from "/components/skeletons/PageSkeleton"; import { PageSkeleton } from "/components/skeletons/PageSkeleton";
import { lazy } from "@probo/react-lazy"; import { lazy } from "@probo/react-lazy";
import { dataQuery, datumNodeQuery } from "../hooks/graph/DatumGraph"; import { dataQuery, datumNodeQuery } from "../hooks/graph/DatumGraph";
@@ -12,7 +12,7 @@ export const dataRoutes = [
path: "data", path: "data",
Fallback: PageSkeleton, Fallback: PageSkeleton,
loader: loaderFromQueryLoader(({ organizationId }) => loader: loaderFromQueryLoader(({ organizationId }) =>
loadQuery<DatumGraphListQuery>(relayEnvironment, dataQuery, { loadQuery<DatumGraphListQuery>(consoleEnvironment, dataQuery, {
organizationId: organizationId, organizationId: organizationId,
snapshotId: null snapshotId: null
}), }),
@@ -25,7 +25,7 @@ export const dataRoutes = [
path: "snapshots/:snapshotId/data", path: "snapshots/:snapshotId/data",
Fallback: PageSkeleton, Fallback: PageSkeleton,
loader: loaderFromQueryLoader(({ organizationId, snapshotId }) => loader: loaderFromQueryLoader(({ organizationId, snapshotId }) =>
loadQuery<DatumGraphListQuery>(relayEnvironment, dataQuery, { loadQuery<DatumGraphListQuery>(consoleEnvironment, dataQuery, {
organizationId, organizationId,
snapshotId, snapshotId,
}), }),
@@ -38,7 +38,7 @@ export const dataRoutes = [
path: "data/:dataId", path: "data/:dataId",
Fallback: PageSkeleton, Fallback: PageSkeleton,
loader: loaderFromQueryLoader(({ dataId }) => loader: loaderFromQueryLoader(({ dataId }) =>
loadQuery<DatumGraphNodeQuery>(relayEnvironment, datumNodeQuery, { dataId }), loadQuery<DatumGraphNodeQuery>(consoleEnvironment, datumNodeQuery, { dataId }),
), ),
Component: withQueryRef(lazy( Component: withQueryRef(lazy(
() => import("../pages/organizations/data/DatumDetailsPage") () => import("../pages/organizations/data/DatumDetailsPage")
@@ -48,7 +48,7 @@ export const dataRoutes = [
path: "snapshots/:snapshotId/data/:dataId", path: "snapshots/:snapshotId/data/:dataId",
Fallback: PageSkeleton, Fallback: PageSkeleton,
loader: loaderFromQueryLoader(({ dataId }) => loader: loaderFromQueryLoader(({ dataId }) =>
loadQuery<DatumGraphNodeQuery>(relayEnvironment, datumNodeQuery, { dataId }), loadQuery<DatumGraphNodeQuery>(consoleEnvironment, datumNodeQuery, { dataId }),
), ),
Component: withQueryRef(lazy( Component: withQueryRef(lazy(
() => import("../pages/organizations/data/DatumDetailsPage") () => import("../pages/organizations/data/DatumDetailsPage")

View File

@@ -1,6 +1,6 @@
import { Fragment } from "react"; import { Fragment } from "react";
import { loadQuery } from "react-relay"; import { loadQuery } from "react-relay";
import { relayEnvironment } from "/providers/RelayProviders"; import { consoleEnvironment } from "/environments";
import { documentsQuery } from "/hooks/graph/DocumentGraph"; import { documentsQuery } from "/hooks/graph/DocumentGraph";
import { documentNodeQuery } from "/hooks/graph/DocumentGraph"; import { documentNodeQuery } from "/hooks/graph/DocumentGraph";
import { PageSkeleton } from "/components/skeletons/PageSkeleton"; import { PageSkeleton } from "/components/skeletons/PageSkeleton";
@@ -58,7 +58,7 @@ export const documentsRoutes = [
path: "documents", path: "documents",
Fallback: PageSkeleton, Fallback: PageSkeleton,
loader: loaderFromQueryLoader(({ organizationId }) => loader: loaderFromQueryLoader(({ organizationId }) =>
loadQuery<DocumentGraphListQuery>(relayEnvironment, documentsQuery, { organizationId: organizationId! }), loadQuery<DocumentGraphListQuery>(consoleEnvironment, documentsQuery, { organizationId: organizationId! }),
), ),
Component: withQueryRef(lazy( Component: withQueryRef(lazy(
() => import("/pages/organizations/documents/DocumentsPage"), () => import("/pages/organizations/documents/DocumentsPage"),
@@ -68,7 +68,7 @@ export const documentsRoutes = [
path: "documents/:documentId", path: "documents/:documentId",
Fallback: PageSkeleton, Fallback: PageSkeleton,
loader: loaderFromQueryLoader(({ documentId }) => loader: loaderFromQueryLoader(({ documentId }) =>
loadQuery<DocumentGraphNodeQuery>(relayEnvironment, documentNodeQuery, { documentId: documentId! }), loadQuery<DocumentGraphNodeQuery>(consoleEnvironment, documentNodeQuery, { documentId: documentId! }),
), ),
Component: withQueryRef(lazy( Component: withQueryRef(lazy(
() => import("../pages/organizations/documents/DocumentDetailPage"), () => import("../pages/organizations/documents/DocumentDetailPage"),

View File

@@ -1,6 +1,6 @@
import { loadQuery } from "react-relay"; import { loadQuery } from "react-relay";
import { PageSkeleton } from "/components/skeletons/PageSkeleton"; import { PageSkeleton } from "/components/skeletons/PageSkeleton";
import { relayEnvironment } from "/providers/RelayProviders"; import { consoleEnvironment } from "/environments";
import { import {
frameworksQuery, frameworksQuery,
frameworkNodeQuery, frameworkNodeQuery,
@@ -19,7 +19,7 @@ export const frameworkRoutes = [
path: "frameworks", path: "frameworks",
Fallback: PageSkeleton, Fallback: PageSkeleton,
loader: loaderFromQueryLoader(({ organizationId }) => loader: loaderFromQueryLoader(({ organizationId }) =>
loadQuery<FrameworkGraphListQuery>(relayEnvironment, frameworksQuery, { organizationId: organizationId! }), loadQuery<FrameworkGraphListQuery>(consoleEnvironment, frameworksQuery, { organizationId: organizationId! }),
), ),
Component: withQueryRef(lazy( Component: withQueryRef(lazy(
() => import("/pages/organizations/frameworks/FrameworksPage") () => import("/pages/organizations/frameworks/FrameworksPage")
@@ -29,7 +29,7 @@ export const frameworkRoutes = [
path: "frameworks/:frameworkId", path: "frameworks/:frameworkId",
Fallback: PageSkeleton, Fallback: PageSkeleton,
loader: loaderFromQueryLoader(({ frameworkId }) => loader: loaderFromQueryLoader(({ frameworkId }) =>
loadQuery<FrameworkGraphNodeQuery>(relayEnvironment, frameworkNodeQuery, { frameworkId: frameworkId! }), loadQuery<FrameworkGraphNodeQuery>(consoleEnvironment, frameworkNodeQuery, { frameworkId: frameworkId! }),
), ),
Component: withQueryRef(lazy( Component: withQueryRef(lazy(
() => import("/pages/organizations/frameworks/FrameworkDetailPage") () => import("/pages/organizations/frameworks/FrameworkDetailPage")
@@ -43,7 +43,7 @@ export const frameworkRoutes = [
path: "controls/:controlId", path: "controls/:controlId",
Fallback: ControlSkeleton, Fallback: ControlSkeleton,
loader: loaderFromQueryLoader(({ controlId }) => loader: loaderFromQueryLoader(({ controlId }) =>
loadQuery<FrameworkGraphControlNodeQuery>(relayEnvironment, frameworkControlNodeQuery, { controlId: controlId! }), loadQuery<FrameworkGraphControlNodeQuery>(consoleEnvironment, frameworkControlNodeQuery, { controlId: controlId! }),
), ),
Component: withQueryRef(lazy( Component: withQueryRef(lazy(
() => import("/pages/organizations/frameworks/FrameworkControlPage") () => import("/pages/organizations/frameworks/FrameworkControlPage")

View File

@@ -1,6 +1,6 @@
import { Fragment } from "react"; import { Fragment } from "react";
import { loadQuery } from "react-relay"; import { loadQuery } from "react-relay";
import { relayEnvironment } from "/providers/RelayProviders"; import { consoleEnvironment } from "/environments";
import { measureNodeQuery, measuresQuery } from "/hooks/graph/MeasureGraph"; import { measureNodeQuery, measuresQuery } from "/hooks/graph/MeasureGraph";
import { PageSkeleton } from "/components/skeletons/PageSkeleton"; import { PageSkeleton } from "/components/skeletons/PageSkeleton";
import { redirect } from "react-router"; import { redirect } from "react-router";
@@ -15,7 +15,7 @@ export const measureRoutes = [
path: "measures", path: "measures",
Fallback: PageSkeleton, Fallback: PageSkeleton,
loader: loaderFromQueryLoader(({ organizationId }) => loader: loaderFromQueryLoader(({ organizationId }) =>
loadQuery<MeasureGraphListQuery>(relayEnvironment, measuresQuery, { organizationId: organizationId! }), loadQuery<MeasureGraphListQuery>(consoleEnvironment, measuresQuery, { organizationId: organizationId! }),
), ),
Component: withQueryRef(lazy(() => import("/pages/organizations/measures/MeasuresPage"))), Component: withQueryRef(lazy(() => import("/pages/organizations/measures/MeasuresPage"))),
children: [ children: [
@@ -29,7 +29,7 @@ export const measureRoutes = [
path: "measures/:measureId", path: "measures/:measureId",
Fallback: PageSkeleton, Fallback: PageSkeleton,
loader: loaderFromQueryLoader(({ measureId }) => loader: loaderFromQueryLoader(({ measureId }) =>
loadQuery<MeasureGraphNodeQuery>(relayEnvironment, measureNodeQuery, { measureId: measureId! }), loadQuery<MeasureGraphNodeQuery>(consoleEnvironment, measureNodeQuery, { measureId: measureId! }),
), ),
Component: withQueryRef(lazy(() => import("/pages/organizations/measures/MeasureDetailPage"))), Component: withQueryRef(lazy(() => import("/pages/organizations/measures/MeasureDetailPage"))),
children: [ children: [

View File

@@ -1,6 +1,6 @@
import { Fragment } from "react"; import { Fragment } from "react";
import { loadQuery } from "react-relay"; import { loadQuery } from "react-relay";
import { relayEnvironment } from "/providers/RelayProviders"; import { consoleEnvironment } from "/environments";
import { meetingsQuery } from "/hooks/graph/MeetingGraph"; import { meetingsQuery } from "/hooks/graph/MeetingGraph";
import { meetingNodeQuery } from "/hooks/graph/MeetingGraph"; import { meetingNodeQuery } from "/hooks/graph/MeetingGraph";
import { PageSkeleton } from "/components/skeletons/PageSkeleton"; import { PageSkeleton } from "/components/skeletons/PageSkeleton";
@@ -40,7 +40,7 @@ export const meetingsRoutes = [
path: "meetings", path: "meetings",
Fallback: PageSkeleton, Fallback: PageSkeleton,
loader: loaderFromQueryLoader(({ organizationId }) => loader: loaderFromQueryLoader(({ organizationId }) =>
loadQuery<MeetingGraphListQuery>(relayEnvironment, meetingsQuery, { organizationId }), loadQuery<MeetingGraphListQuery>(consoleEnvironment, meetingsQuery, { organizationId }),
), ),
Component: withQueryRef(lazy( Component: withQueryRef(lazy(
() => import("/pages/organizations/meetings/MeetingsPage"), () => import("/pages/organizations/meetings/MeetingsPage"),
@@ -50,7 +50,7 @@ export const meetingsRoutes = [
path: "meetings/:meetingId", path: "meetings/:meetingId",
Fallback: PageSkeleton, Fallback: PageSkeleton,
loader: loaderFromQueryLoader(({ meetingId }) => loader: loaderFromQueryLoader(({ meetingId }) =>
loadQuery<MeetingGraphNodeQuery>(relayEnvironment, meetingNodeQuery, { meetingId }), loadQuery<MeetingGraphNodeQuery>(consoleEnvironment, meetingNodeQuery, { meetingId }),
), ),
Component: withQueryRef(lazy( Component: withQueryRef(lazy(
() => import("../pages/organizations/meetings/MeetingDetailPage"), () => import("../pages/organizations/meetings/MeetingDetailPage"),

View File

@@ -1,5 +1,5 @@
import { loadQuery } from "react-relay"; import { loadQuery } from "react-relay";
import { relayEnvironment } from "/providers/RelayProviders"; import { consoleEnvironment } from "/environments";
import { PageSkeleton } from "/components/skeletons/PageSkeleton"; import { PageSkeleton } from "/components/skeletons/PageSkeleton";
import { lazy } from "@probo/react-lazy"; import { lazy } from "@probo/react-lazy";
import { nonconformitiesQuery, nonconformityNodeQuery } from "../hooks/graph/NonconformityGraph"; import { nonconformitiesQuery, nonconformityNodeQuery } from "../hooks/graph/NonconformityGraph";
@@ -12,7 +12,7 @@ export const nonconformityRoutes = [
path: "nonconformities", path: "nonconformities",
Fallback: PageSkeleton, Fallback: PageSkeleton,
loader: loaderFromQueryLoader(({ organizationId }) => loader: loaderFromQueryLoader(({ organizationId }) =>
loadQuery<NonconformityGraphListQuery>(relayEnvironment, nonconformitiesQuery, { loadQuery<NonconformityGraphListQuery>(consoleEnvironment, nonconformitiesQuery, {
organizationId, organizationId,
snapshotId: null snapshotId: null
}), }),
@@ -25,7 +25,7 @@ export const nonconformityRoutes = [
path: "snapshots/:snapshotId/nonconformities", path: "snapshots/:snapshotId/nonconformities",
Fallback: PageSkeleton, Fallback: PageSkeleton,
loader: loaderFromQueryLoader(({ organizationId, snapshotId }) => loader: loaderFromQueryLoader(({ organizationId, snapshotId }) =>
loadQuery<NonconformityGraphListQuery>(relayEnvironment, nonconformitiesQuery, { loadQuery<NonconformityGraphListQuery>(consoleEnvironment, nonconformitiesQuery, {
organizationId, organizationId,
snapshotId, snapshotId,
}), }),
@@ -38,7 +38,7 @@ export const nonconformityRoutes = [
path: "nonconformities/:nonconformityId", path: "nonconformities/:nonconformityId",
Fallback: PageSkeleton, Fallback: PageSkeleton,
loader: loaderFromQueryLoader(({ nonconformityId }) => loader: loaderFromQueryLoader(({ nonconformityId }) =>
loadQuery<NonconformityGraphNodeQuery>(relayEnvironment, nonconformityNodeQuery, { loadQuery<NonconformityGraphNodeQuery>(consoleEnvironment, nonconformityNodeQuery, {
nonconformityId, nonconformityId,
}), }),
), ),
@@ -50,7 +50,7 @@ export const nonconformityRoutes = [
path: "snapshots/:snapshotId/nonconformities/:nonconformityId", path: "snapshots/:snapshotId/nonconformities/:nonconformityId",
Fallback: PageSkeleton, Fallback: PageSkeleton,
loader: loaderFromQueryLoader(({ nonconformityId }) => loader: loaderFromQueryLoader(({ nonconformityId }) =>
loadQuery<NonconformityGraphNodeQuery>(relayEnvironment, nonconformityNodeQuery, { loadQuery<NonconformityGraphNodeQuery>(consoleEnvironment, nonconformityNodeQuery, {
nonconformityId: nonconformityId nonconformityId: nonconformityId
}), }),
), ),

View File

@@ -1,5 +1,5 @@
import { loadQuery } from "react-relay"; import { loadQuery } from "react-relay";
import { relayEnvironment } from "/providers/RelayProviders"; import { consoleEnvironment } from "/environments";
import { PageSkeleton } from "/components/skeletons/PageSkeleton"; import { PageSkeleton } from "/components/skeletons/PageSkeleton";
import { lazy } from "@probo/react-lazy"; import { lazy } from "@probo/react-lazy";
import { obligationsQuery, obligationNodeQuery } from "/hooks/graph/ObligationGraph"; import { obligationsQuery, obligationNodeQuery } from "/hooks/graph/ObligationGraph";
@@ -12,7 +12,7 @@ export const obligationRoutes = [
path: "obligations", path: "obligations",
Fallback: PageSkeleton, Fallback: PageSkeleton,
loader: loaderFromQueryLoader(({ organizationId }) => loader: loaderFromQueryLoader(({ organizationId }) =>
loadQuery<ObligationGraphListQuery>(relayEnvironment, obligationsQuery, { loadQuery<ObligationGraphListQuery>(consoleEnvironment, obligationsQuery, {
organizationId, organizationId,
snapshotId: null snapshotId: null
}), }),
@@ -25,7 +25,7 @@ export const obligationRoutes = [
path: "snapshots/:snapshotId/obligations", path: "snapshots/:snapshotId/obligations",
Fallback: PageSkeleton, Fallback: PageSkeleton,
loader: loaderFromQueryLoader(({ organizationId, snapshotId }) => loader: loaderFromQueryLoader(({ organizationId, snapshotId }) =>
loadQuery<ObligationGraphListQuery>(relayEnvironment, obligationsQuery, { loadQuery<ObligationGraphListQuery>(consoleEnvironment, obligationsQuery, {
organizationId, organizationId,
snapshotId, snapshotId,
}), }),
@@ -38,7 +38,7 @@ export const obligationRoutes = [
path: "obligations/:obligationId", path: "obligations/:obligationId",
Fallback: PageSkeleton, Fallback: PageSkeleton,
loader: loaderFromQueryLoader(({ obligationId }) => loader: loaderFromQueryLoader(({ obligationId }) =>
loadQuery<ObligationGraphNodeQuery>(relayEnvironment, obligationNodeQuery, { loadQuery<ObligationGraphNodeQuery>(consoleEnvironment, obligationNodeQuery, {
obligationId, obligationId,
}), }),
), ),
@@ -50,7 +50,7 @@ export const obligationRoutes = [
path: "snapshots/:snapshotId/obligations/:obligationId", path: "snapshots/:snapshotId/obligations/:obligationId",
Fallback: PageSkeleton, Fallback: PageSkeleton,
loader: loaderFromQueryLoader(({ obligationId }) => loader: loaderFromQueryLoader(({ obligationId }) =>
loadQuery<ObligationGraphNodeQuery>(relayEnvironment, obligationNodeQuery, { loadQuery<ObligationGraphNodeQuery>(consoleEnvironment, obligationNodeQuery, {
obligationId, obligationId,
}), }),
), ),

View File

@@ -1,6 +1,6 @@
import { lazy } from "@probo/react-lazy"; import { lazy } from "@probo/react-lazy";
import { loadQuery } from "react-relay"; import { loadQuery } from "react-relay";
import { relayEnvironment } from "/providers/RelayProviders"; import { consoleEnvironment } from "/environments";
import { PageSkeleton } from "/components/skeletons/PageSkeleton.tsx"; import { PageSkeleton } from "/components/skeletons/PageSkeleton.tsx";
import { import {
paginatedPeopleQuery, paginatedPeopleQuery,
@@ -16,7 +16,7 @@ export const peopleRoutes = [
path: "people", path: "people",
Fallback: PageSkeleton, Fallback: PageSkeleton,
loader: loaderFromQueryLoader(({ organizationId }) => loader: loaderFromQueryLoader(({ organizationId }) =>
loadQuery<PeopleGraphPaginatedQuery>(relayEnvironment, paginatedPeopleQuery, { organizationId: organizationId! }), loadQuery<PeopleGraphPaginatedQuery>(consoleEnvironment, paginatedPeopleQuery, { organizationId: organizationId! }),
), ),
Component: withQueryRef(lazy(() => import("/pages/organizations/people/PeopleListPage"))), Component: withQueryRef(lazy(() => import("/pages/organizations/people/PeopleListPage"))),
}, },
@@ -24,7 +24,7 @@ export const peopleRoutes = [
path: "people/:peopleId", path: "people/:peopleId",
Fallback: PageSkeleton, Fallback: PageSkeleton,
loader: loaderFromQueryLoader(({ peopleId }) => loader: loaderFromQueryLoader(({ peopleId }) =>
loadQuery<PeopleGraphNodeQuery>(relayEnvironment, peopleNodeQuery, { peopleId: peopleId! }), loadQuery<PeopleGraphNodeQuery>(consoleEnvironment, peopleNodeQuery, { peopleId: peopleId! }),
), ),
Component: withQueryRef(lazy( Component: withQueryRef(lazy(
() => import("/pages/organizations/people/PeopleDetailPage") () => import("/pages/organizations/people/PeopleDetailPage")

View File

@@ -1,5 +1,5 @@
import { loadQuery } from "react-relay"; import { loadQuery } from "react-relay";
import { relayEnvironment } from "/providers/RelayProviders"; import { consoleEnvironment } from "/environments";
import { PageSkeleton } from "/components/skeletons/PageSkeleton"; import { PageSkeleton } from "/components/skeletons/PageSkeleton";
import { lazy } from "@probo/react-lazy"; import { lazy } from "@probo/react-lazy";
import { processingActivitiesQuery, processingActivityNodeQuery } from "/hooks/graph/ProcessingActivityGraph"; import { processingActivitiesQuery, processingActivityNodeQuery } from "/hooks/graph/ProcessingActivityGraph";
@@ -12,7 +12,7 @@ export const processingActivityRoutes = [
path: "processing-activities", path: "processing-activities",
Fallback: PageSkeleton, Fallback: PageSkeleton,
loader: loaderFromQueryLoader(({ organizationId }) => loader: loaderFromQueryLoader(({ organizationId }) =>
loadQuery<ProcessingActivityGraphListQuery>(relayEnvironment, processingActivitiesQuery, { loadQuery<ProcessingActivityGraphListQuery>(consoleEnvironment, processingActivitiesQuery, {
organizationId: organizationId!, organizationId: organizationId!,
snapshotId: null snapshotId: null
}), }),
@@ -25,7 +25,7 @@ export const processingActivityRoutes = [
path: "snapshots/:snapshotId/processing-activities", path: "snapshots/:snapshotId/processing-activities",
Fallback: PageSkeleton, Fallback: PageSkeleton,
loader: loaderFromQueryLoader(({ organizationId, snapshotId }) => loader: loaderFromQueryLoader(({ organizationId, snapshotId }) =>
loadQuery<ProcessingActivityGraphListQuery>(relayEnvironment, processingActivitiesQuery, { loadQuery<ProcessingActivityGraphListQuery>(consoleEnvironment, processingActivitiesQuery, {
organizationId: organizationId!, organizationId: organizationId!,
snapshotId, snapshotId,
}), }),
@@ -38,7 +38,7 @@ export const processingActivityRoutes = [
path: "processing-activities/:activityId", path: "processing-activities/:activityId",
Fallback: PageSkeleton, Fallback: PageSkeleton,
loader: loaderFromQueryLoader(({ activityId }) => loader: loaderFromQueryLoader(({ activityId }) =>
loadQuery<ProcessingActivityGraphNodeQuery>(relayEnvironment, processingActivityNodeQuery, { loadQuery<ProcessingActivityGraphNodeQuery>(consoleEnvironment, processingActivityNodeQuery, {
processingActivityId: activityId!, processingActivityId: activityId!,
}), }),
), ),
@@ -50,7 +50,7 @@ export const processingActivityRoutes = [
path: "snapshots/:snapshotId/processing-activities/:activityId", path: "snapshots/:snapshotId/processing-activities/:activityId",
Fallback: PageSkeleton, Fallback: PageSkeleton,
loader: loaderFromQueryLoader(({ activityId }) => loader: loaderFromQueryLoader(({ activityId }) =>
loadQuery<ProcessingActivityGraphNodeQuery>(relayEnvironment, processingActivityNodeQuery, { loadQuery<ProcessingActivityGraphNodeQuery>(consoleEnvironment, processingActivityNodeQuery, {
processingActivityId: activityId!, processingActivityId: activityId!,
}), }),
), ),

View File

@@ -1,7 +1,7 @@
import { Fragment } from "react"; import { Fragment } from "react";
import { loadQuery } from "react-relay"; import { loadQuery } from "react-relay";
import { RisksPageSkeleton } from "/components/skeletons/RisksPageSkeleton.tsx"; import { RisksPageSkeleton } from "/components/skeletons/RisksPageSkeleton.tsx";
import { relayEnvironment } from "/providers/RelayProviders"; import { consoleEnvironment } from "/environments";
import { riskNodeQuery, risksQuery } from "/hooks/graph/RiskGraph"; import { riskNodeQuery, risksQuery } from "/hooks/graph/RiskGraph";
import { PageSkeleton } from "/components/skeletons/PageSkeleton"; import { PageSkeleton } from "/components/skeletons/PageSkeleton";
import { redirect } from "react-router"; import { redirect } from "react-router";
@@ -16,7 +16,7 @@ export const riskRoutes = [
path: "risks", path: "risks",
Fallback: RisksPageSkeleton, Fallback: RisksPageSkeleton,
loader: loaderFromQueryLoader(({ organizationId }) => loader: loaderFromQueryLoader(({ organizationId }) =>
loadQuery<RiskGraphListQuery>(relayEnvironment, risksQuery, { loadQuery<RiskGraphListQuery>(consoleEnvironment, risksQuery, {
organizationId: organizationId!, organizationId: organizationId!,
snapshotId: null snapshotId: null
}), }),
@@ -27,7 +27,7 @@ export const riskRoutes = [
path: "snapshots/:snapshotId/risks", path: "snapshots/:snapshotId/risks",
Fallback: RisksPageSkeleton, Fallback: RisksPageSkeleton,
loader: loaderFromQueryLoader(({ organizationId, snapshotId }) => loader: loaderFromQueryLoader(({ organizationId, snapshotId }) =>
loadQuery<RiskGraphListQuery>(relayEnvironment, risksQuery, { loadQuery<RiskGraphListQuery>(consoleEnvironment, risksQuery, {
organizationId: organizationId!, organizationId: organizationId!,
snapshotId: snapshotId! snapshotId: snapshotId!
}), }),
@@ -38,7 +38,7 @@ export const riskRoutes = [
path: "risks/:riskId", path: "risks/:riskId",
Fallback: PageSkeleton, Fallback: PageSkeleton,
loader: loaderFromQueryLoader(({ riskId }) => loader: loaderFromQueryLoader(({ riskId }) =>
loadQuery<RiskGraphNodeQuery>(relayEnvironment, riskNodeQuery, { riskId: riskId! }), loadQuery<RiskGraphNodeQuery>(consoleEnvironment, riskNodeQuery, { riskId: riskId! }),
), ),
Component: withQueryRef(lazy(() => import("/pages/organizations/risks/RiskDetailPage"))), Component: withQueryRef(lazy(() => import("/pages/organizations/risks/RiskDetailPage"))),
children: [ children: [
@@ -90,7 +90,7 @@ export const riskRoutes = [
path: "snapshots/:snapshotId/risks/:riskId", path: "snapshots/:snapshotId/risks/:riskId",
Fallback: PageSkeleton, Fallback: PageSkeleton,
loader: loaderFromQueryLoader(({ riskId }) => loader: loaderFromQueryLoader(({ riskId }) =>
loadQuery<RiskGraphNodeQuery>(relayEnvironment, riskNodeQuery, { riskId: riskId! }), loadQuery<RiskGraphNodeQuery>(consoleEnvironment, riskNodeQuery, { riskId: riskId! }),
), ),
Component: withQueryRef(lazy(() => import("/pages/organizations/risks/RiskDetailPage"))), Component: withQueryRef(lazy(() => import("/pages/organizations/risks/RiskDetailPage"))),
children: [ children: [

View File

@@ -1,6 +1,6 @@
import { loadQuery } from "react-relay"; import { loadQuery } from "react-relay";
import { PageSkeleton } from "/components/skeletons/PageSkeleton"; import { PageSkeleton } from "/components/skeletons/PageSkeleton";
import { relayEnvironment } from "/providers/RelayProviders"; import { consoleEnvironment } from "/environments";
import { snapshotsQuery, snapshotNodeQuery } from "/hooks/graph/SnapshotGraph"; import { snapshotsQuery, snapshotNodeQuery } from "/hooks/graph/SnapshotGraph";
import { lazy } from "@probo/react-lazy"; import { lazy } from "@probo/react-lazy";
import type { SnapshotGraphListQuery } from "/hooks/graph/__generated__/SnapshotGraphListQuery.graphql"; import type { SnapshotGraphListQuery } from "/hooks/graph/__generated__/SnapshotGraphListQuery.graphql";
@@ -12,7 +12,7 @@ export const snapshotsRoutes = [
path: "snapshots", path: "snapshots",
Fallback: PageSkeleton, Fallback: PageSkeleton,
loader: loaderFromQueryLoader(({ organizationId }) => loader: loaderFromQueryLoader(({ organizationId }) =>
loadQuery<SnapshotGraphListQuery>(relayEnvironment, snapshotsQuery, { organizationId }), loadQuery<SnapshotGraphListQuery>(consoleEnvironment, snapshotsQuery, { organizationId }),
), ),
Component: withQueryRef(lazy( Component: withQueryRef(lazy(
() => import("/pages/organizations/snapshots/SnapshotsPage") () => import("/pages/organizations/snapshots/SnapshotsPage")
@@ -22,7 +22,7 @@ export const snapshotsRoutes = [
path: "snapshots/:snapshotId", path: "snapshots/:snapshotId",
Fallback: PageSkeleton, Fallback: PageSkeleton,
loader: loaderFromQueryLoader(({ snapshotId }) => loader: loaderFromQueryLoader(({ snapshotId }) =>
loadQuery<SnapshotGraphNodeQuery>(relayEnvironment, snapshotNodeQuery, { snapshotId }), loadQuery<SnapshotGraphNodeQuery>(consoleEnvironment, snapshotNodeQuery, { snapshotId }),
), ),
Component: withQueryRef(lazy( Component: withQueryRef(lazy(
() => import("/pages/organizations/snapshots/SnapshotDetailPage") () => import("/pages/organizations/snapshots/SnapshotDetailPage")

View File

@@ -1,7 +1,7 @@
import { lazy } from "@probo/react-lazy"; import { lazy } from "@probo/react-lazy";
import { PageSkeleton } from "/components/skeletons/PageSkeleton"; import { PageSkeleton } from "/components/skeletons/PageSkeleton";
import { loadQuery } from "react-relay"; import { loadQuery } from "react-relay";
import { relayEnvironment } from "/providers/RelayProviders"; import { consoleEnvironment } from "/environments";
import { tasksQuery } from "/hooks/graph/TaskGraph"; import { tasksQuery } from "/hooks/graph/TaskGraph";
import { loaderFromQueryLoader, withQueryRef, type AppRoute } from "@probo/routes"; import { loaderFromQueryLoader, withQueryRef, type AppRoute } from "@probo/routes";
import type { TaskGraphQuery } from "/hooks/graph/__generated__/TaskGraphQuery.graphql"; import type { TaskGraphQuery } from "/hooks/graph/__generated__/TaskGraphQuery.graphql";
@@ -10,7 +10,7 @@ export const taskRoutes = [
path: "tasks", path: "tasks",
Fallback: PageSkeleton, Fallback: PageSkeleton,
loader: loaderFromQueryLoader(({ organizationId }) => loader: loaderFromQueryLoader(({ organizationId }) =>
loadQuery<TaskGraphQuery>(relayEnvironment, tasksQuery, { loadQuery<TaskGraphQuery>(consoleEnvironment, tasksQuery, {
organizationId, organizationId,
}), }),
), ),

View File

@@ -1,5 +1,5 @@
import { loadQuery } from "react-relay"; import { loadQuery } from "react-relay";
import { relayEnvironment } from "/providers/RelayProviders"; import { consoleEnvironment } from "/environments";
import { PageSkeleton } from "/components/skeletons/PageSkeleton"; import { PageSkeleton } from "/components/skeletons/PageSkeleton";
import { LinkCardSkeleton } from "/components/skeletons/LinkCardSkeleton"; import { LinkCardSkeleton } from "/components/skeletons/LinkCardSkeleton";
import { lazy } from "@probo/react-lazy"; import { lazy } from "@probo/react-lazy";
@@ -12,7 +12,7 @@ export const trustCenterRoutes = [
path: "trust-center", path: "trust-center",
Fallback: PageSkeleton, Fallback: PageSkeleton,
loader: loaderFromQueryLoader(({ organizationId }) => loader: loaderFromQueryLoader(({ organizationId }) =>
loadQuery<TrustCenterGraphQuery>(relayEnvironment, trustCenterQuery, { organizationId }, { fetchPolicy: "network-only" }), loadQuery<TrustCenterGraphQuery>(consoleEnvironment, trustCenterQuery, { organizationId }, { fetchPolicy: "network-only" }),
), ),
Component: withQueryRef(lazy( Component: withQueryRef(lazy(
() => import("/pages/organizations/trustCenter/TrustCenterPage") () => import("/pages/organizations/trustCenter/TrustCenterPage")

View File

@@ -1,6 +1,6 @@
import { lazy } from "@probo/react-lazy"; import { lazy } from "@probo/react-lazy";
import { loadQuery } from "react-relay"; import { loadQuery } from "react-relay";
import { relayEnvironment } from "/providers/RelayProviders"; import { consoleEnvironment } from "/environments";
import { PageSkeleton } from "/components/skeletons/PageSkeleton"; import { PageSkeleton } from "/components/skeletons/PageSkeleton";
import { vendorNodeQuery, vendorsQuery } from "/hooks/graph/VendorGraph"; import { vendorNodeQuery, vendorsQuery } from "/hooks/graph/VendorGraph";
import { LinkCardSkeleton } from "/components/skeletons/LinkCardSkeleton"; import { LinkCardSkeleton } from "/components/skeletons/LinkCardSkeleton";
@@ -13,7 +13,7 @@ export const vendorRoutes = [
path: "vendors", path: "vendors",
Fallback: PageSkeleton, Fallback: PageSkeleton,
loader: loaderFromQueryLoader(({ organizationId }) => loader: loaderFromQueryLoader(({ organizationId }) =>
loadQuery<VendorGraphListQuery>(relayEnvironment, vendorsQuery, { loadQuery<VendorGraphListQuery>(consoleEnvironment, vendorsQuery, {
organizationId: organizationId!, organizationId: organizationId!,
snapshotId: null snapshotId: null
}), }),
@@ -24,7 +24,7 @@ export const vendorRoutes = [
path: "snapshots/:snapshotId/vendors", path: "snapshots/:snapshotId/vendors",
Fallback: PageSkeleton, Fallback: PageSkeleton,
loader: loaderFromQueryLoader(({ organizationId, snapshotId }) => loader: loaderFromQueryLoader(({ organizationId, snapshotId }) =>
loadQuery<VendorGraphListQuery>(relayEnvironment, vendorsQuery, { loadQuery<VendorGraphListQuery>(consoleEnvironment, vendorsQuery, {
organizationId: organizationId!, organizationId: organizationId!,
snapshotId snapshotId
}), }),
@@ -35,7 +35,7 @@ export const vendorRoutes = [
path: "vendors/:vendorId", path: "vendors/:vendorId",
Fallback: PageSkeleton, Fallback: PageSkeleton,
loader: loaderFromQueryLoader(({ vendorId }) => loader: loaderFromQueryLoader(({ vendorId }) =>
loadQuery<VendorGraphNodeQuery>(relayEnvironment, vendorNodeQuery, { loadQuery<VendorGraphNodeQuery>(consoleEnvironment, vendorNodeQuery, {
vendorId: vendorId!, vendorId: vendorId!,
}), }),
), ),
@@ -100,7 +100,7 @@ export const vendorRoutes = [
path: "snapshots/:snapshotId/vendors/:vendorId", path: "snapshots/:snapshotId/vendors/:vendorId",
Fallback: PageSkeleton, Fallback: PageSkeleton,
loader: loaderFromQueryLoader(({ vendorId }) => loader: loaderFromQueryLoader(({ vendorId }) =>
loadQuery<VendorGraphNodeQuery>(relayEnvironment, vendorNodeQuery, { loadQuery<VendorGraphNodeQuery>(consoleEnvironment, vendorNodeQuery, {
vendorId: vendorId!, vendorId: vendorId!,
}), }),
), ),

View File

@@ -24,6 +24,7 @@ export default defineConfig({
}, },
resolve: { resolve: {
alias: { alias: {
"/environments": fileURLToPath(new URL("./src/environments", import.meta.url)),
"/type": fileURLToPath(new URL("./src/type.ts", import.meta.url)), "/type": fileURLToPath(new URL("./src/type.ts", import.meta.url)),
"/components": fileURLToPath( "/components": fileURLToPath(
new URL("./src/components", import.meta.url), new URL("./src/components", import.meta.url),

View File

@@ -66,7 +66,7 @@ const fetchRelay: FetchFunction = async (
request, request,
variables, variables,
_, _,
uploadables, uploadables
) => { ) => {
const requestInit: RequestInit = { const requestInit: RequestInit = {
method: "POST", method: "POST",
@@ -82,7 +82,7 @@ const fetchRelay: FetchFunction = async (
operationName: request.name, operationName: request.name,
query: request.text, query: request.text,
variables: variables, variables: variables,
}), })
); );
const uploadableMap: { const uploadableMap: {
@@ -124,7 +124,7 @@ const fetchRelay: FetchFunction = async (
// For custom domains at /overview, this resolves to /api/trust/v1/graphql // For custom domains at /overview, this resolves to /api/trust/v1/graphql
const response = await fetch( const response = await fetch(
buildEndpoint("./api/trust/v1/graphql"), buildEndpoint("./api/trust/v1/graphql"),
requestInit, requestInit
); );
if (response.status === 500) { if (response.status === 500) {
@@ -161,7 +161,7 @@ const store = new Store(source, {
gcReleaseBufferSize: 20, gcReleaseBufferSize: 20,
}); });
export const relayEnvironment = new Environment({ export const consoleEnvironment = new Environment({
network: Network.create(fetchRelay), network: Network.create(fetchRelay),
store, store,
}); });
@@ -171,7 +171,7 @@ export const relayEnvironment = new Environment({
*/ */
export function RelayProvider({ children }: PropsWithChildren) { export function RelayProvider({ children }: PropsWithChildren) {
return ( return (
<RelayEnvironmentProvider environment={relayEnvironment}> <RelayEnvironmentProvider environment={consoleEnvironment}>
{children} {children}
</RelayEnvironmentProvider> </RelayEnvironmentProvider>
); );

View File

@@ -6,7 +6,7 @@ import {
} from "react-router"; } from "react-router";
import { Fragment } from "react"; import { Fragment } from "react";
import { import {
relayEnvironment, consoleEnvironment,
UnAuthenticatedError, UnAuthenticatedError,
} from "./providers/RelayProviders.tsx"; } from "./providers/RelayProviders.tsx";
import { loadQuery } from "react-relay"; import { loadQuery } from "react-relay";
@@ -23,7 +23,12 @@ import { SubprocessorsPage } from "/pages/SubprocessorsPage";
import { AccessPage } from "./pages/AccessPage.tsx"; import { AccessPage } from "./pages/AccessPage.tsx";
import { TabSkeleton } from "./components/Skeletons/TabSkeleton.tsx"; import { TabSkeleton } from "./components/Skeletons/TabSkeleton.tsx";
import { MainSkeleton } from "./components/Skeletons/MainSkeleton.tsx"; import { MainSkeleton } from "./components/Skeletons/MainSkeleton.tsx";
import { loaderFromQueryLoader, routeFromAppRoute, withQueryRef, type AppRoute } from "@probo/routes"; import {
loaderFromQueryLoader,
routeFromAppRoute,
withQueryRef,
type AppRoute,
} from "@probo/routes";
/** /**
* Top level error boundary * Top level error boundary
@@ -50,7 +55,9 @@ const routes = [
// Custom domain routes (subdomain-based) // Custom domain routes (subdomain-based)
{ {
path: "/overview", path: "/overview",
loader: loaderFromQueryLoader(() => loadQuery(relayEnvironment, currentTrustGraphQuery, {})), loader: loaderFromQueryLoader(() =>
loadQuery(consoleEnvironment, currentTrustGraphQuery, {})
),
Component: withQueryRef(MainLayout), Component: withQueryRef(MainLayout),
Fallback: MainSkeleton, Fallback: MainSkeleton,
ErrorBoundary: ErrorBoundary, ErrorBoundary: ErrorBoundary,
@@ -64,14 +71,18 @@ const routes = [
}, },
{ {
path: "/documents", path: "/documents",
loader: loaderFromQueryLoader(() => loadQuery(relayEnvironment, currentTrustGraphQuery, {})), loader: loaderFromQueryLoader(() =>
loadQuery(consoleEnvironment, currentTrustGraphQuery, {})
),
Component: withQueryRef(MainLayout), Component: withQueryRef(MainLayout),
Fallback: MainSkeleton, Fallback: MainSkeleton,
ErrorBoundary: ErrorBoundary, ErrorBoundary: ErrorBoundary,
children: [ children: [
{ {
path: "", path: "",
loader: loaderFromQueryLoader(() => loadQuery(relayEnvironment, currentTrustDocumentsQuery, {})), loader: loaderFromQueryLoader(() =>
loadQuery(consoleEnvironment, currentTrustDocumentsQuery, {})
),
Fallback: TabSkeleton, Fallback: TabSkeleton,
Component: withQueryRef(DocumentsPage), Component: withQueryRef(DocumentsPage),
}, },
@@ -79,14 +90,18 @@ const routes = [
}, },
{ {
path: "/subprocessors", path: "/subprocessors",
loader: loaderFromQueryLoader(() => loadQuery(relayEnvironment, currentTrustGraphQuery, {})), loader: loaderFromQueryLoader(() =>
loadQuery(consoleEnvironment, currentTrustGraphQuery, {})
),
Component: withQueryRef(MainLayout), Component: withQueryRef(MainLayout),
Fallback: MainSkeleton, Fallback: MainSkeleton,
ErrorBoundary: ErrorBoundary, ErrorBoundary: ErrorBoundary,
children: [ children: [
{ {
path: "", path: "",
loader: loaderFromQueryLoader(() => loadQuery(relayEnvironment, currentTrustVendorsQuery, {})), loader: loaderFromQueryLoader(() =>
loadQuery(consoleEnvironment, currentTrustVendorsQuery, {})
),
Fallback: TabSkeleton, Fallback: TabSkeleton,
Component: withQueryRef(SubprocessorsPage), Component: withQueryRef(SubprocessorsPage),
}, },

View File

@@ -13,15 +13,36 @@ probod:
encryption-key: "thisisnotasecretAAAAAAAAAAAAAAAAAAAAAAAAAAA=" encryption-key: "thisisnotasecretAAAAAAAAAAAAAAAAAAAAAAAAAAA="
chrome-dp-addr: "localhost:9222" chrome-dp-addr: "localhost:9222"
identity-and-access-management:
signup:
enabled: true
invitation-token-validity: 3600
password:
pepper: "this-is-a-secure-pepper-for-password-hashing-at-least-32-bytes"
session:
duration: 7d
cookie:
name: "SSID"
domain: "localhost"
secret: "this-is-a-secure-secret-for-cookie-signing-at-least-32-bytes"
duration: 24
secure: true
trust-center: trust-center:
http-addr: ":8085" http-addr: ":8085"
https-addr: ":8443" https-addr: ":8443"
api: api:
addr: "localhost:8080" addr: "localhost:8080"
cors:
allowed-origins: ["http://localhost:8080", "http://localhost:5173"] default:
extra-header-fields: {} cors:
allowed-origins: ["http://localhost:8080", "http://localhost:5173"]
extra-header-fields: {}
console:
complaince-page:
connect:
mcp:
pg: pg:
addr: "localhost:5432" addr: "localhost:5432"

View File

@@ -146,7 +146,6 @@ services:
command: command:
- "start-dev" - "start-dev"
- "--import-realm" - "--import-realm"
- "--verbose"
ports: ports:
- 8082:8080 - 8082:8080
volumes: volumes:

4
go.sum
View File

@@ -2,14 +2,10 @@ codeberg.org/miekg/dns v0.6.2 h1:gkuKad3tKCs4On/3ZlydASeXPc1uQaZNuDSxEXR7Mt0=
codeberg.org/miekg/dns v0.6.2/go.mod h1:IKSpRNHVdUyxHC457VnDQ9Dv05UhXXsALZAywCM7D54= codeberg.org/miekg/dns v0.6.2/go.mod h1:IKSpRNHVdUyxHC457VnDQ9Dv05UhXXsALZAywCM7D54=
github.com/99designs/gqlgen v0.17.84 h1:iVMdiStgUVx/BFkMb0J5GAXlqfqtQ7bqMCYK6v52kQ0= github.com/99designs/gqlgen v0.17.84 h1:iVMdiStgUVx/BFkMb0J5GAXlqfqtQ7bqMCYK6v52kQ0=
github.com/99designs/gqlgen v0.17.84/go.mod h1:qjoUqzTeiejdo+bwUg8unqSpeYG42XrcrQboGIezmFA= github.com/99designs/gqlgen v0.17.84/go.mod h1:qjoUqzTeiejdo+bwUg8unqSpeYG42XrcrQboGIezmFA=
github.com/PuerkitoBio/goquery v1.11.0 h1:jZ7pwMQXIITcUXNH83LLk+txlaEy6NVOfTuP43xxfqw=
github.com/PuerkitoBio/goquery v1.11.0/go.mod h1:wQHgxUOU3JGuj3oD/QFfxUdlzW6xPHfqyHre6VMY4DQ=
github.com/agnivade/levenshtein v1.2.1 h1:EHBY3UOn1gwdy/VbFwgo4cxecRznFk7fKWN1KOX7eoM= github.com/agnivade/levenshtein v1.2.1 h1:EHBY3UOn1gwdy/VbFwgo4cxecRznFk7fKWN1KOX7eoM=
github.com/agnivade/levenshtein v1.2.1/go.mod h1:QVVI16kDrtSuwcpd0p1+xMC6Z/VfhtCyDIjcwga4/DU= github.com/agnivade/levenshtein v1.2.1/go.mod h1:QVVI16kDrtSuwcpd0p1+xMC6Z/VfhtCyDIjcwga4/DU=
github.com/andreyvit/diff v0.0.0-20170406064948-c7f18ee00883 h1:bvNMNQO63//z+xNgfBlViaCIJKLlCJ6/fmUseuG0wVQ= github.com/andreyvit/diff v0.0.0-20170406064948-c7f18ee00883 h1:bvNMNQO63//z+xNgfBlViaCIJKLlCJ6/fmUseuG0wVQ=
github.com/andreyvit/diff v0.0.0-20170406064948-c7f18ee00883/go.mod h1:rCTlJbsFo29Kk6CurOXKm700vrz8f0KW0JNfpkRJY/8= github.com/andreyvit/diff v0.0.0-20170406064948-c7f18ee00883/go.mod h1:rCTlJbsFo29Kk6CurOXKm700vrz8f0KW0JNfpkRJY/8=
github.com/andybalholm/cascadia v1.3.3 h1:AG2YHrzJIm4BZ19iwJ/DAua6Btl3IwJX+VI4kktS1LM=
github.com/andybalholm/cascadia v1.3.3/go.mod h1:xNd9bqTn98Ln4DwST8/nG+H0yuB8Hmgu1YHNnWw0GeA=
github.com/arbovm/levenshtein v0.0.0-20160628152529-48b4e1c0c4d0 h1:jfIu9sQUG6Ig+0+Ap1h4unLjW6YQJpKZVmUzxsD4E/Q= github.com/arbovm/levenshtein v0.0.0-20160628152529-48b4e1c0c4d0 h1:jfIu9sQUG6Ig+0+Ap1h4unLjW6YQJpKZVmUzxsD4E/Q=
github.com/arbovm/levenshtein v0.0.0-20160628152529-48b4e1c0c4d0/go.mod h1:t2tdKJDJF9BV14lnkjHmOQgcvEKgtqs5a1N3LNdJhGE= github.com/arbovm/levenshtein v0.0.0-20160628152529-48b4e1c0c4d0/go.mod h1:t2tdKJDJF9BV14lnkjHmOQgcvEKgtqs5a1N3LNdJhGE=
github.com/aws/aws-sdk-go-v2 v1.40.0 h1:/WMUA0kjhZExjOQN2z3oLALDREea1A7TobfuiBrKlwc= github.com/aws/aws-sdk-go-v2 v1.40.0 h1:/WMUA0kjhZExjOQN2z3oLALDREea1A7TobfuiBrKlwc=

1969
package-lock.json generated

File diff suppressed because it is too large Load Diff

View File

@@ -29,8 +29,5 @@
"packageManager": "npm@11.1.0", "packageManager": "npm@11.1.0",
"volta": { "volta": {
"node": "24.4.0" "node": "24.4.0"
},
"dependencies": {
"react-dom": "^19.2.1"
} }
} }

View File

@@ -11,6 +11,7 @@
"dependencies": { "dependencies": {
"@react-email/components": "^0.5.6", "@react-email/components": "^0.5.6",
"react": "^19.0.0", "react": "^19.0.0",
"react-dom": "^19.0.0",
"react-email": "^4.3.0" "react-email": "^4.3.0"
}, },
"devDependencies": { "devDependencies": {

View File

@@ -0,0 +1,31 @@
{
"name": "@probo/relay",
"version": "1.0.0",
"description": "",
"type": "module",
"main": "./src/index.ts",
"scripts": {
"test": "vitest run",
"test:watch": "vitest",
"format": "prettier --write ./src"
},
"prettier": "@probo/prettier",
"keywords": [],
"author": "",
"license": "ISC",
"devDependencies": {
"@probo/prettier": "1.0.0",
"prettier": "^3.5.3",
"typescript": "^5.8.3",
"vitest": "^3.1.3"
},
"peerDependencies": {
"react": "^19.2.1",
"react-relay": "^19.0.0",
"react-router": "^7.10.0",
"relay-runtime": "^19.0.0"
},
"dependencies": {
"graphql": "^16.12.0"
}
}

View File

@@ -0,0 +1,53 @@
export class UnAuthenticatedError extends Error {
constructor(message?: string) {
super(message || "UNAUTHENTICATED");
this.name = "UnAuthenticatedError";
Object.setPrototypeOf(this, UnAuthenticatedError.prototype);
}
}
export class InternalServerError extends Error {
constructor() {
super("INTERNAL_SERVER_ERROR");
this.name = "InternalServerError";
Object.setPrototypeOf(this, InternalServerError.prototype);
}
}
export class AuthenticationRequiredError extends Error {
public redirectUrl: string;
public requiresSaml: boolean;
public organizationId: string;
public samlConfigId?: string;
constructor(extensions: {
redirectUrl: string;
requiresSaml: boolean;
organizationId: string;
samlConfigId?: string;
}) {
super("AUTHENTICATION_REQUIRED");
this.name = "AuthenticationRequiredError";
Object.setPrototypeOf(this, AuthenticationRequiredError.prototype);
this.redirectUrl = extensions.redirectUrl;
this.requiresSaml = extensions.requiresSaml;
this.organizationId = extensions.organizationId;
this.samlConfigId = extensions.samlConfigId;
}
}
export class UnauthorizedError extends Error {
constructor(message?: string) {
super(message || "UNAUTHORIZED");
this.name = "UnauthorizedError";
Object.setPrototypeOf(this, UnauthorizedError.prototype);
}
}
export class ForbiddenError extends Error {
constructor(message?: string) {
super(message || "FORBIDDEN");
this.name = "ForbiddenError";
Object.setPrototypeOf(this, ForbiddenError.prototype);
}
}

112
packages/relay/src/fetch.ts Normal file
View File

@@ -0,0 +1,112 @@
import { type FetchFunction } from "relay-runtime";
import {
InternalServerError,
UnAuthenticatedError,
AuthenticationRequiredError,
UnauthorizedError,
ForbiddenError,
} from "./errors";
import { GraphQLError } from "graphql";
const hasUnauthenticatedError = (error: GraphQLError) =>
error.extensions?.code == "UNAUTHENTICATED";
const hasAuthenticationRequiredError = (error: GraphQLError) =>
error.extensions?.code == "AUTHENTICATION_REQUIRED";
const hasUnauthorizedError = (error: GraphQLError) =>
error.extensions?.code == "UNAUTHORIZED";
const hasForbiddenError = (error: GraphQLError) =>
error.extensions?.code == "FORBIDDEN";
export const makeFetchQuery = (endpoint: string): FetchFunction => {
return async (request, variables, _, uploadables) => {
const requestInit: RequestInit = {
method: "POST",
credentials: "include",
headers: {},
};
if (uploadables) {
const formData = new FormData();
formData.append(
"operations",
JSON.stringify({
operationName: request.name,
query: request.text,
variables: variables,
}),
);
const uploadableMap: {
[key: string]: string[];
} = {};
Object.keys(uploadables).forEach((key, index) => {
uploadableMap[index] = [`variables.${key}`];
});
formData.append("map", JSON.stringify(uploadableMap));
Object.keys(uploadables).forEach((key, index) => {
formData.append(index.toString(), uploadables[key]);
});
requestInit.body = formData;
} else {
requestInit.headers = {
Accept: "application/graphql-response+json; charset=utf-8, application/json; charset=utf-8",
"Content-Type": "application/json",
};
requestInit.body = JSON.stringify({
operationName: request.name,
query: request.text,
variables,
});
}
const response = await fetch(endpoint, requestInit);
if (response.status === 500) {
throw new InternalServerError();
}
const json = await response.json();
if (json.errors) {
const errors = json.errors as GraphQLError[];
const unauthenticatedError = errors.find(hasUnauthenticatedError);
if (unauthenticatedError) {
throw new UnAuthenticatedError(unauthenticatedError.message);
}
const authRequiredError = errors.find(hasAuthenticationRequiredError);
if (authRequiredError?.extensions) {
const { redirectUrl, requiresSaml, organizationId, samlConfigId } =
authRequiredError.extensions;
throw new AuthenticationRequiredError({
redirectUrl: redirectUrl as string,
requiresSaml: requiresSaml as boolean,
organizationId: organizationId as string,
samlConfigId: samlConfigId as string | undefined,
});
}
const unauthorizedError = errors.find(hasUnauthorizedError);
if (unauthorizedError) {
throw new UnauthorizedError(unauthorizedError.message);
}
const forbiddenError = errors.find(hasForbiddenError);
if (forbiddenError) {
throw new ForbiddenError(forbiddenError.message);
}
}
return json;
};
};

View File

@@ -0,0 +1,2 @@
export { makeFetchQuery } from "./fetch";
export * from "./errors";

View File

@@ -0,0 +1,6 @@
{
"compilerOptions": {
"lib": ["ES2021", "dom"],
"jsx": "react-jsx",
}
}

View File

@@ -1,88 +0,0 @@
package auth
import (
"fmt"
"go.probo.inc/probo/pkg/coredata"
"go.probo.inc/probo/pkg/gid"
)
// AuthMethod represents a method of authentication
type AuthMethod int
const (
AuthMethodPassword AuthMethod = iota
AuthMethodSAML
AuthMethodAny
)
type OrgAuthRequirement struct {
OrganizationID gid.GID
EmailDomain string
SAMLConfig *coredata.SAMLConfiguration
}
type AccessResult struct {
OrganizationID gid.GID
Allowed bool
MissingAuth AuthMethod
SAMLConfig *coredata.SAMLConfiguration
}
func (r OrgAuthRequirement) Check(session coredata.SessionData) AccessResult {
if r.SAMLConfig == nil || !r.SAMLConfig.Enabled || !r.SAMLConfig.DomainVerified {
return AccessResult{
OrganizationID: r.OrganizationID,
Allowed: session.PasswordAuthenticated,
MissingAuth: AuthMethodPassword,
SAMLConfig: nil,
}
}
orgKey := r.OrganizationID.String()
_, hasSAML := session.SAMLAuthenticatedOrgs[orgKey]
if r.SAMLConfig.EnforcementPolicy == coredata.SAMLEnforcementPolicyRequired {
return AccessResult{
OrganizationID: r.OrganizationID,
Allowed: hasSAML,
MissingAuth: AuthMethodSAML,
SAMLConfig: r.SAMLConfig,
}
}
hasAnyAuth := session.PasswordAuthenticated || hasSAML
missingAuth := AuthMethodAny
if hasAnyAuth {
missingAuth = AuthMethodPassword
}
return AccessResult{
OrganizationID: r.OrganizationID,
Allowed: hasAnyAuth,
MissingAuth: missingAuth,
SAMLConfig: r.SAMLConfig,
}
}
func (r AccessResult) ToError(baseURL string) error {
if r.Allowed {
return nil
}
switch r.MissingAuth {
case AuthMethodPassword:
return ErrPasswordAuthRequired{
OrganizationID: r.OrganizationID,
RedirectURL: fmt.Sprintf("%s/auth/login?method=password", baseURL),
}
case AuthMethodSAML, AuthMethodAny:
return ErrSAMLAuthRequired{
ConfigID: r.SAMLConfig.ID,
OrganizationID: r.OrganizationID,
RedirectURL: fmt.Sprintf("%s/connect/saml/login/%s", baseURL, r.SAMLConfig.ID),
}
default:
return fmt.Errorf("access denied to organization %s", r.OrganizationID)
}
}

View File

@@ -1,109 +0,0 @@
// Copyright (c) 2025 Probo Inc <hello@getprobo.com>.
//
// Permission to use, copy, modify, and/or distribute this software for any
// purpose with or without fee is hereby granted, provided that the above
// copyright notice and this permission notice appear in all copies.
//
// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH
// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY
// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT,
// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM
// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR
// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR
// PERFORMANCE OF THIS SOFTWARE.
package auth
import (
"context"
"time"
"go.gearno.de/kit/log"
"go.gearno.de/kit/pg"
)
const (
DefaultCleanupInterval = 1 * time.Hour
)
type (
Cleaner struct {
pg *pg.Client
interval time.Duration
logger *log.Logger
}
)
func NewCleaner(
pg *pg.Client,
interval time.Duration,
logger *log.Logger,
) *Cleaner {
if interval == 0 {
interval = DefaultCleanupInterval
}
return &Cleaner{
pg: pg,
interval: interval,
logger: logger.Named("saml.cleaner"),
}
}
func (c *Cleaner) Run(ctx context.Context) error {
c.logger.InfoCtx(ctx, "SAML cleaner starting", log.Duration("interval", c.interval))
if err := c.cleanup(ctx); err != nil {
c.logger.ErrorCtx(ctx, "initial cleanup failed", log.Error(err))
}
ticker := time.NewTicker(c.interval)
defer ticker.Stop()
for {
select {
case <-ctx.Done():
c.logger.InfoCtx(ctx, "SAML cleaner shutting down")
return ctx.Err()
case <-ticker.C:
if err := c.cleanup(ctx); err != nil {
c.logger.ErrorCtx(ctx, "periodic cleanup failed", log.Error(err))
}
}
}
}
func (c *Cleaner) cleanup(ctx context.Context) error {
var assertionsDeleted, requestsDeleted int64
err := c.pg.WithConn(
ctx,
func(conn pg.Conn) error {
count, err := CleanupExpiredAssertions(ctx, conn)
if err != nil {
return err
}
assertionsDeleted = count
count, err = CleanupExpiredRequests(ctx, conn)
if err != nil {
return err
}
requestsDeleted = count
return nil
},
)
if err != nil {
return err
}
if assertionsDeleted > 0 || requestsDeleted > 0 {
c.logger.InfoCtx(ctx, "cleaned up expired SAML data",
log.Int64("assertions", assertionsDeleted),
log.Int64("requests", requestsDeleted))
}
return nil
}

View File

@@ -1,288 +0,0 @@
// Copyright (c) 2025 Probo Inc <hello@getprobo.com>.
//
// Permission to use, copy, modify, and/or distribute this software for any
// purpose with or without fee is hereby granted, provided that the above
// copyright notice and this permission notice appear in all copies.
//
// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH
// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY
// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT,
// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM
// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR
// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR
// PERFORMANCE OF THIS SOFTWARE.
package auth
import (
"context"
"fmt"
"strings"
"time"
"go.probo.inc/probo/pkg/coredata"
"go.probo.inc/probo/pkg/gid"
"go.gearno.de/kit/pg"
)
type (
CreateSAMLConfigurationRequest struct {
OrganizationID gid.GID
EmailDomain string
EnforcementPolicy coredata.SAMLEnforcementPolicy
IdPEntityID string
IdPSsoURL string
IdPCertificate string
IdPMetadataURL *string
AttributeEmail string
AttributeFirstname string
AttributeLastname string
AttributeRole string
AutoSignupEnabled bool
}
UpdateSAMLConfigurationRequest struct {
ID gid.GID
Enabled *bool
EnforcementPolicy *coredata.SAMLEnforcementPolicy
IdPEntityID *string
IdPSsoURL *string
IdPCertificate *string
IdPMetadataURL *string
AttributeEmail *string
AttributeFirstname *string
AttributeLastname *string
AttributeRole *string
AutoSignupEnabled *bool
}
)
func (s TenantAuthService) CreateSAMLConfiguration(
ctx context.Context,
req CreateSAMLConfigurationRequest,
) (*coredata.SAMLConfiguration, error) {
// Validate only the IdP configuration (user-provided data)
if err := ValidateIdPConfiguration(req.IdPEntityID, req.IdPSsoURL, req.IdPCertificate); err != nil {
return nil, fmt.Errorf("SAML configuration validation failed: %w", err)
}
var config *coredata.SAMLConfiguration
err := s.pg.WithTx(
ctx,
func(tx pg.Conn) error {
now := time.Now()
tenantID := s.scope.GetTenantID()
var org coredata.Organization
if err := org.LoadByID(ctx, tx, s.scope, req.OrganizationID); err != nil {
return fmt.Errorf("organization not found: %w", err)
}
config = &coredata.SAMLConfiguration{
ID: gid.New(tenantID, coredata.SAMLConfigurationEntityType),
OrganizationID: org.ID,
EmailDomain: req.EmailDomain,
EnforcementPolicy: req.EnforcementPolicy,
Enabled: false,
IdPEntityID: req.IdPEntityID,
IdPSsoURL: req.IdPSsoURL,
IdPCertificate: req.IdPCertificate,
IdPMetadataURL: req.IdPMetadataURL,
AttributeEmail: req.AttributeEmail,
AttributeFirstname: req.AttributeFirstname,
AttributeLastname: req.AttributeLastname,
AttributeRole: req.AttributeRole,
AutoSignupEnabled: req.AutoSignupEnabled,
CreatedAt: now,
UpdatedAt: now,
}
if err := config.Insert(ctx, tx, s.scope); err != nil {
return fmt.Errorf("cannot insert saml configuration: %w", err)
}
return nil
},
)
if err != nil {
return nil, err
}
return config, nil
}
func (s TenantAuthService) UpdateSAMLConfiguration(
ctx context.Context,
req UpdateSAMLConfigurationRequest,
) (*coredata.SAMLConfiguration, error) {
var config *coredata.SAMLConfiguration
err := s.pg.WithTx(
ctx,
func(tx pg.Conn) error {
var cfg coredata.SAMLConfiguration
if err := cfg.LoadByID(ctx, tx, s.scope, req.ID); err != nil {
return fmt.Errorf("cannot load saml configuration: %w", err)
}
if req.Enabled != nil {
cfg.Enabled = *req.Enabled
}
if req.EnforcementPolicy != nil {
cfg.EnforcementPolicy = *req.EnforcementPolicy
}
if req.IdPEntityID != nil {
cfg.IdPEntityID = *req.IdPEntityID
}
if req.IdPSsoURL != nil {
cfg.IdPSsoURL = *req.IdPSsoURL
}
if req.IdPCertificate != nil {
cfg.IdPCertificate = *req.IdPCertificate
}
if req.IdPMetadataURL != nil {
cfg.IdPMetadataURL = req.IdPMetadataURL
}
if req.AttributeEmail != nil {
cfg.AttributeEmail = *req.AttributeEmail
}
if req.AttributeFirstname != nil {
cfg.AttributeFirstname = *req.AttributeFirstname
}
if req.AttributeLastname != nil {
cfg.AttributeLastname = *req.AttributeLastname
}
if req.AttributeRole != nil {
cfg.AttributeRole = *req.AttributeRole
}
if req.AutoSignupEnabled != nil {
cfg.AutoSignupEnabled = *req.AutoSignupEnabled
}
cfg.UpdatedAt = time.Now()
if err := cfg.Update(ctx, tx, s.scope); err != nil {
return fmt.Errorf("cannot update saml configuration: %w", err)
}
config = &cfg
return nil
},
)
if err != nil {
return nil, err
}
return config, nil
}
func (s TenantAuthService) DeleteSAMLConfiguration(
ctx context.Context,
configID gid.GID,
) error {
return s.pg.WithTx(
ctx,
func(tx pg.Conn) error {
var config coredata.SAMLConfiguration
if err := config.LoadByID(ctx, tx, s.scope, configID); err != nil {
return fmt.Errorf("cannot load saml configuration: %w", err)
}
if err := config.Delete(ctx, tx, s.scope); err != nil {
return fmt.Errorf("cannot delete saml configuration: %w", err)
}
return nil
},
)
}
func (s TenantAuthService) EnableSAMLConfiguration(
ctx context.Context,
configID gid.GID,
) (*coredata.SAMLConfiguration, error) {
enabled := true
return s.UpdateSAMLConfiguration(ctx, UpdateSAMLConfigurationRequest{
ID: configID,
Enabled: &enabled,
})
}
func (s TenantAuthService) DisableSAMLConfiguration(
ctx context.Context,
configID gid.GID,
) (*coredata.SAMLConfiguration, error) {
disabled := false
return s.UpdateSAMLConfiguration(ctx, UpdateSAMLConfigurationRequest{
ID: configID,
Enabled: &disabled,
})
}
func (s TenantAuthService) GetSAMLConfigurationByID(
ctx context.Context,
configID gid.GID,
) (*coredata.SAMLConfiguration, error) {
var config coredata.SAMLConfiguration
err := s.pg.WithConn(
ctx,
func(conn pg.Conn) error {
return config.LoadByID(ctx, conn, s.scope, configID)
},
)
if err != nil {
return nil, fmt.Errorf("cannot load saml configuration: %w", err)
}
return &config, nil
}
func (s TenantAuthService) GetSAMLConfigurationsByOrganizationID(
ctx context.Context,
organizationID gid.GID,
) ([]*coredata.SAMLConfiguration, error) {
var configs []*coredata.SAMLConfiguration
err := s.pg.WithConn(
ctx,
func(conn pg.Conn) error {
var err error
configs, err = coredata.LoadSAMLConfigurationsByOrganizationID(ctx, conn, s.scope, organizationID)
return err
},
)
if err != nil {
return nil, fmt.Errorf("cannot load saml configurations: %w", err)
}
return configs, nil
}
func (s Service) CheckSSOAvailabilityByEmail(
ctx context.Context,
email string,
) ([]*coredata.SAMLConfiguration, error) {
// Extract domain from email
parts := strings.Split(email, "@")
if len(parts) != 2 {
return nil, fmt.Errorf("invalid email format")
}
domain := parts[1]
var configs []*coredata.SAMLConfiguration
err := s.pg.WithConn(
ctx,
func(conn pg.Conn) error {
var err error
configs, err = coredata.LoadAllEnabledSAMLConfigurationsByEmailDomain(ctx, conn, domain)
return err
},
)
if err != nil {
return nil, fmt.Errorf("cannot load saml configurations: %w", err)
}
return configs, nil
}

View File

@@ -1,204 +0,0 @@
// Copyright (c) 2025 Probo Inc <hello@getprobo.com>.
//
// Permission to use, copy, modify, and/or distribute this software for any
// purpose with or without fee is hereby granted, provided that the above
// copyright notice and this permission notice appear in all copies.
//
// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH
// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY
// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT,
// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM
// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR
// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR
// PERFORMANCE OF THIS SOFTWARE.
package auth
import (
"crypto/rand"
"crypto/rsa"
"crypto/x509"
"crypto/x509/pkix"
"encoding/base64"
"encoding/pem"
"encoding/xml"
"fmt"
"math/big"
"time"
"github.com/crewjam/saml"
)
func GenerateServiceProviderMetadata(
entityID string,
acsURL string,
spCert *x509.Certificate,
) ([]byte, error) {
certData := base64.StdEncoding.EncodeToString(spCert.Raw)
trueVal := true
metadata := &saml.EntityDescriptor{
EntityID: entityID,
SPSSODescriptors: []saml.SPSSODescriptor{
{
SSODescriptor: saml.SSODescriptor{
RoleDescriptor: saml.RoleDescriptor{
ProtocolSupportEnumeration: "urn:oasis:names:tc:SAML:2.0:protocol",
KeyDescriptors: []saml.KeyDescriptor{
{
Use: "signing",
KeyInfo: saml.KeyInfo{
X509Data: saml.X509Data{
X509Certificates: []saml.X509Certificate{
{Data: certData},
},
},
},
},
{
Use: "encryption",
KeyInfo: saml.KeyInfo{
X509Data: saml.X509Data{
X509Certificates: []saml.X509Certificate{
{Data: certData},
},
},
},
},
},
},
},
AuthnRequestsSigned: &trueVal,
WantAssertionsSigned: &trueVal,
AssertionConsumerServices: []saml.IndexedEndpoint{
{
Binding: saml.HTTPPostBinding,
Location: acsURL,
Index: 0,
},
},
},
},
}
xmlBytes, err := xml.MarshalIndent(metadata, "", " ")
if err != nil {
return nil, fmt.Errorf("cannot marshal SP metadata to XML: %w", err)
}
return xmlBytes, nil
}
func ParseIdPCertificate(certPEM string) (*x509.Certificate, error) {
block, _ := pem.Decode([]byte(certPEM))
if block == nil {
return nil, fmt.Errorf("cannot decode PEM block from IdP certificate")
}
cert, err := x509.ParseCertificate(block.Bytes)
if err != nil {
return nil, fmt.Errorf("cannot parse X.509 certificate: %w", err)
}
return cert, nil
}
type IdPMetadata struct {
EntityID string
SsoURL string
Certificate string
MetadataURL *string
}
func ParseIdPMetadata(metadataXML string) (*IdPMetadata, error) {
var entityDescriptor saml.EntityDescriptor
if err := xml.Unmarshal([]byte(metadataXML), &entityDescriptor); err != nil {
return nil, fmt.Errorf("cannot parse IdP metadata XML: %w", err)
}
if len(entityDescriptor.IDPSSODescriptors) == 0 {
return nil, fmt.Errorf("no IDPSSODescriptor found in metadata")
}
idpDescriptor := entityDescriptor.IDPSSODescriptors[0]
var ssoURL string
for _, sso := range idpDescriptor.SingleSignOnServices {
if sso.Binding == saml.HTTPPostBinding || sso.Binding == saml.HTTPRedirectBinding {
ssoURL = sso.Location
break
}
}
if ssoURL == "" && len(idpDescriptor.SingleSignOnServices) > 0 {
ssoURL = idpDescriptor.SingleSignOnServices[0].Location
}
if ssoURL == "" {
return nil, fmt.Errorf("no SingleSignOnService found in metadata")
}
var certPEM string
for _, keyDescriptor := range idpDescriptor.KeyDescriptors {
if keyDescriptor.Use == "signing" || keyDescriptor.Use == "" {
if len(keyDescriptor.KeyInfo.X509Data.X509Certificates) > 0 {
certData := keyDescriptor.KeyInfo.X509Data.X509Certificates[0].Data
certDER, err := base64.StdEncoding.DecodeString(certData)
if err != nil {
return nil, fmt.Errorf("cannot decode certificate: %w", err)
}
certPEM = string(pem.EncodeToMemory(&pem.Block{
Type: "CERTIFICATE",
Bytes: certDER,
}))
break
}
}
}
if certPEM == "" {
return nil, fmt.Errorf("no signing certificate found in metadata")
}
return &IdPMetadata{
EntityID: entityDescriptor.EntityID,
SsoURL: ssoURL,
Certificate: certPEM,
}, nil
}
func GenerateSelfSignedCertificate(entityID string) (*x509.Certificate, *rsa.PrivateKey, error) {
privateKey, err := rsa.GenerateKey(rand.Reader, 2048)
if err != nil {
return nil, nil, fmt.Errorf("cannot generate RSA private key: %w", err)
}
serialNumberLimit := new(big.Int).Lsh(big.NewInt(1), 128)
serialNumber, err := rand.Int(rand.Reader, serialNumberLimit)
if err != nil {
return nil, nil, fmt.Errorf("cannot generate serial number: %w", err)
}
template := x509.Certificate{
SerialNumber: serialNumber,
Subject: pkix.Name{
CommonName: entityID,
Organization: []string{"Probo"},
},
NotBefore: time.Now(),
NotAfter: time.Now().Add(10 * 365 * 24 * time.Hour),
KeyUsage: x509.KeyUsageKeyEncipherment | x509.KeyUsageDigitalSignature,
ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth},
BasicConstraintsValid: true,
}
certDER, err := x509.CreateCertificate(rand.Reader, &template, &template, &privateKey.PublicKey, privateKey)
if err != nil {
return nil, nil, fmt.Errorf("cannot create certificate: %w", err)
}
cert, err := x509.ParseCertificate(certDER)
if err != nil {
return nil, nil, fmt.Errorf("cannot parse created certificate: %w", err)
}
return cert, privateKey, nil
}

View File

@@ -1,556 +0,0 @@
// Copyright (c) 2025 Probo Inc <hello@getprobo.com>.
//
// Permission to use, copy, modify, and/or distribute this software for any
// purpose with or without fee is hereby granted, provided that the above
// copyright notice and this permission notice appear in all copies.
//
// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH
// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY
// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT,
// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM
// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR
// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR
// PERFORMANCE OF THIS SOFTWARE.
package auth
import (
"context"
"crypto/rsa"
"crypto/x509"
"encoding/base64"
"encoding/pem"
"errors"
"fmt"
"net/http"
"net/url"
"strings"
"time"
"github.com/crewjam/saml"
"go.gearno.de/kit/log"
"go.gearno.de/kit/pg"
"go.probo.inc/probo/pkg/coredata"
"go.probo.inc/probo/pkg/crypto/cipher"
"go.probo.inc/probo/pkg/gid"
"go.probo.inc/probo/pkg/mail"
)
type (
SAMLService struct {
pg *pg.Client
encryptionKey cipher.EncryptionKey
baseURL string
sessionDuration time.Duration
cookieName string
cookieSecret string
certificate *x509.Certificate
privateKey *rsa.PrivateKey
logger *log.Logger
}
ErrSPCertificateNotConfigured struct{}
ErrSAMLConfigurationNotFound struct {
OrganizationID gid.GID
}
ErrSAMLDisabled struct {
OrganizationID gid.GID
}
ErrInvalidIdPCertificate struct {
Err error
}
ErrInvalidURL struct {
Field string
URL string
Err error
}
ErrCannotCreateServiceProvider struct {
Err error
}
ErrCannotCreateAuthRequest struct {
Err error
}
ErrCannotGenerateRedirectURL struct {
Err error
}
ErrCannotParseSAMLResponse struct {
Err error
}
ErrCannotValidateAssertion struct {
Err error
}
ErrCannotExtractUserAttributes struct {
Err error
}
ErrCannotMapRole struct {
Err error
}
ErrReplayAttackDetected struct {
AssertionID string
Err error
}
)
func (e ErrSPCertificateNotConfigured) Error() string {
return "SP certificate and private key are not configured"
}
func (e ErrSAMLConfigurationNotFound) Error() string {
return fmt.Sprintf("SAML configuration not found for organization %s", e.OrganizationID)
}
func (e ErrSAMLDisabled) Error() string {
return fmt.Sprintf("SAML is disabled for organization %s", e.OrganizationID)
}
func (e ErrInvalidIdPCertificate) Error() string {
return fmt.Sprintf("cannot parse IdP certificate: %v", e.Err)
}
func (e ErrInvalidURL) Error() string {
return fmt.Sprintf("cannot parse %s URL %q: %v", e.Field, e.URL, e.Err)
}
func (e ErrCannotCreateServiceProvider) Error() string {
return fmt.Sprintf("cannot create service provider: %v", e.Err)
}
func (e ErrCannotCreateAuthRequest) Error() string {
return fmt.Sprintf("cannot create AuthnRequest: %v", e.Err)
}
func (e ErrCannotGenerateRedirectURL) Error() string {
return fmt.Sprintf("cannot generate redirect URL: %v", e.Err)
}
func (e ErrCannotParseSAMLResponse) Error() string {
return fmt.Sprintf("cannot parse SAML response: %v", e.Err)
}
func (e ErrCannotValidateAssertion) Error() string {
return fmt.Sprintf("cannot validate assertion: %v", e.Err)
}
func (e ErrCannotExtractUserAttributes) Error() string {
return fmt.Sprintf("cannot extract user attributes: %v", e.Err)
}
func (e ErrCannotMapRole) Error() string {
return fmt.Sprintf("cannot map role: %v", e.Err)
}
func (e ErrReplayAttackDetected) Error() string {
return fmt.Sprintf("replay attack detected for assertion %s: %v", e.AssertionID, e.Err)
}
func NewSAMLService(
pg *pg.Client,
encryptionKey cipher.EncryptionKey,
baseURL string,
sessionDuration time.Duration,
cookieName string,
cookieSecret string,
certificatePEM string,
privateKeyPEM string,
logger *log.Logger,
) (*SAMLService, error) {
var certificate *x509.Certificate
var privateKey *rsa.PrivateKey
if certificatePEM != "" {
block, _ := pem.Decode([]byte(certificatePEM))
if block == nil || block.Type != "CERTIFICATE" {
return nil, fmt.Errorf("invalid certificate PEM format")
}
cert, err := x509.ParseCertificate(block.Bytes)
if err != nil {
return nil, fmt.Errorf("cannot parse certificate: %w", err)
}
certificate = cert
}
if privateKeyPEM != "" {
block, _ := pem.Decode([]byte(privateKeyPEM))
if block == nil {
return nil, fmt.Errorf("invalid private key PEM format")
}
var key *rsa.PrivateKey
var err error
switch block.Type {
case "RSA PRIVATE KEY":
key, err = x509.ParsePKCS1PrivateKey(block.Bytes)
if err != nil {
return nil, fmt.Errorf("cannot parse PKCS1 private key: %w", err)
}
case "PRIVATE KEY":
parsedKey, err := x509.ParsePKCS8PrivateKey(block.Bytes)
if err != nil {
return nil, fmt.Errorf("cannot parse PKCS8 private key: %w", err)
}
var ok bool
key, ok = parsedKey.(*rsa.PrivateKey)
if !ok {
return nil, fmt.Errorf("private key is not RSA")
}
default:
return nil, fmt.Errorf("unsupported private key type: %s", block.Type)
}
privateKey = key
}
return &SAMLService{
pg: pg,
encryptionKey: encryptionKey,
baseURL: baseURL,
sessionDuration: sessionDuration,
cookieName: cookieName,
cookieSecret: cookieSecret,
certificate: certificate,
privateKey: privateKey,
logger: logger,
}, nil
}
func (s *SAMLService) GetEntityID() string {
return fmt.Sprintf("%s/connect/saml/metadata", s.baseURL)
}
func (s *SAMLService) GetAcsURL() string {
return fmt.Sprintf("%s/connect/saml/consume", s.baseURL)
}
func (s *SAMLService) GetServiceProvider(
ctx context.Context,
config *coredata.SAMLConfiguration,
) (*saml.ServiceProvider, error) {
if s.certificate == nil || s.privateKey == nil {
return nil, ErrSPCertificateNotConfigured{}
}
idpCert, err := ParseIdPCertificate(config.IdPCertificate)
if err != nil {
return nil, ErrInvalidIdPCertificate{Err: err}
}
acsURL, err := url.Parse(s.GetAcsURL())
if err != nil {
return nil, ErrInvalidURL{Field: "ACS", URL: s.GetAcsURL(), Err: err}
}
metadataURL, err := url.Parse(s.GetEntityID())
if err != nil {
return nil, ErrInvalidURL{Field: "Metadata", URL: s.GetEntityID(), Err: err}
}
idpSSOURL, err := url.Parse(config.IdPSsoURL)
if err != nil {
return nil, ErrInvalidURL{Field: "IdP SSO", URL: config.IdPSsoURL, Err: err}
}
sp := &saml.ServiceProvider{
EntityID: s.GetEntityID(),
Key: s.privateKey,
Certificate: s.certificate,
MetadataURL: *metadataURL,
AcsURL: *acsURL,
SloURL: *acsURL,
AllowIDPInitiated: true,
IDPMetadata: &saml.EntityDescriptor{
EntityID: config.IdPEntityID,
IDPSSODescriptors: []saml.IDPSSODescriptor{
{
SSODescriptor: saml.SSODescriptor{
RoleDescriptor: saml.RoleDescriptor{
ProtocolSupportEnumeration: "urn:oasis:names:tc:SAML:2.0:protocol",
KeyDescriptors: []saml.KeyDescriptor{
{
Use: "signing",
KeyInfo: saml.KeyInfo{
X509Data: saml.X509Data{
X509Certificates: []saml.X509Certificate{
{Data: base64.StdEncoding.EncodeToString(idpCert.Raw)},
},
},
},
},
},
},
},
SingleSignOnServices: []saml.Endpoint{
{
Binding: saml.HTTPRedirectBinding,
Location: idpSSOURL.String(),
},
},
},
},
},
}
return sp, nil
}
func (s *SAMLService) InitiateSAMLLogin(
ctx context.Context,
organizationID gid.GID,
tenantID gid.TenantID,
emailDomain string,
) (string, error) {
var config coredata.SAMLConfiguration
scope := coredata.NewScope(tenantID)
err := s.pg.WithConn(
ctx,
func(conn pg.Conn) error {
return config.LoadByOrganizationIDAndEmailDomain(ctx, conn, scope, organizationID, emailDomain)
},
)
if err != nil {
return "", ErrSAMLConfigurationNotFound{OrganizationID: organizationID}
}
if !config.Enabled {
return "", ErrSAMLDisabled{OrganizationID: organizationID}
}
sp, err := s.GetServiceProvider(ctx, &config)
if err != nil {
return "", ErrCannotCreateServiceProvider{Err: err}
}
authReq, err := sp.MakeAuthenticationRequest(
config.IdPSsoURL,
saml.HTTPRedirectBinding,
saml.HTTPPostBinding,
)
if err != nil {
return "", ErrCannotCreateAuthRequest{Err: err}
}
now := time.Now()
requestExpiry := now.Add(10 * time.Minute)
err = s.pg.WithTx(
ctx,
func(tx pg.Conn) error {
samlRequest := coredata.SAMLRequest{
ID: authReq.ID,
OrganizationID: organizationID,
CreatedAt: now,
ExpiresAt: requestExpiry,
}
if err := samlRequest.Insert(ctx, tx, scope); err != nil {
return fmt.Errorf("cannot insert SAML request: %w", err)
}
return nil
},
)
if err != nil {
return "", err
}
redirectURL, err := authReq.Redirect(config.ID.String(), sp)
if err != nil {
return "", ErrCannotGenerateRedirectURL{Err: err}
}
return redirectURL.String(), nil
}
type SAMLUserInfo struct {
Email mail.Addr
FullName string
Role *coredata.MembershipRole
SAMLSubject string
OrganizationID gid.GID
SAMLConfigID gid.GID
}
func (s *SAMLService) loadConfigFromRelayState(
ctx context.Context,
relayStateValue string,
) (*coredata.SAMLConfiguration, *coredata.Organization, error) {
if relayStateValue == "" {
return nil, nil, fmt.Errorf("RelayState is required and must contain SAML config ID")
}
samlConfigID, err := gid.ParseGID(relayStateValue)
if err != nil {
return nil, nil, fmt.Errorf("invalid SAML config ID in RelayState: %w", err)
}
var config coredata.SAMLConfiguration
var org coredata.Organization
err = s.pg.WithConn(
ctx,
func(conn pg.Conn) error {
if err := config.LoadByID(ctx, conn, coredata.NewNoScope(), samlConfigID); err != nil {
return fmt.Errorf("cannot load SAML configuration: %w", err)
}
if err := org.LoadByID(ctx, conn, coredata.NewNoScope(), config.OrganizationID); err != nil {
return fmt.Errorf("organization not found: %w", err)
}
return nil
},
)
if err != nil {
return nil, nil, err
}
return &config, &org, nil
}
func (s *SAMLService) HandleSAMLAssertion(
ctx context.Context,
req *http.Request,
) (*SAMLUserInfo, error) {
samlResponseEncoded := req.FormValue("SAMLResponse")
if samlResponseEncoded == "" {
return nil, fmt.Errorf("missing SAMLResponse in request")
}
relayStateValue := req.FormValue("RelayState")
config, org, err := s.loadConfigFromRelayState(ctx, relayStateValue)
if err != nil {
return nil, err
}
if !config.Enabled {
return nil, ErrSAMLDisabled{OrganizationID: config.OrganizationID}
}
sp, err := s.GetServiceProvider(ctx, config)
if err != nil {
return nil, ErrCannotCreateServiceProvider{Err: err}
}
if req.URL.Scheme == "" {
req.URL.Scheme = "https"
}
if req.URL.Host == "" {
req.URL.Host = req.Host
}
now := time.Now()
var possibleRequestIDs []string
err = s.pg.WithConn(
ctx,
func(conn pg.Conn) error {
requestIDs, err := coredata.LoadValidRequestIDsForOrganization(ctx, conn, config.OrganizationID, now)
if err != nil {
return err
}
possibleRequestIDs = requestIDs
return nil
},
)
if err != nil {
return nil, fmt.Errorf("cannot load valid request IDs: %w", err)
}
assertion, err := sp.ParseResponse(req, possibleRequestIDs)
if err != nil {
return nil, fmt.Errorf("cannot parse SAML response: %w", err)
}
if err := ValidateAssertion(assertion, s.GetEntityID(), now); err != nil {
return nil, ErrCannotValidateAssertion{Err: err}
}
if assertion.ID != "" {
var expiresAt time.Time
if assertion.Conditions != nil && !assertion.Conditions.NotOnOrAfter.IsZero() {
expiresAt = assertion.Conditions.NotOnOrAfter
} else {
expiresAt = now.Add(24 * time.Hour)
}
scope := coredata.NewScope(org.TenantID)
err = s.pg.WithTx(
ctx,
func(tx pg.Conn) error {
if err := PreventReplayAttack(ctx, tx, scope, assertion.ID, config.OrganizationID, expiresAt); err != nil {
return fmt.Errorf("cannot prevent replay attack: %w", err)
}
return nil
},
)
if err != nil {
var replayAttackErr *coredata.ErrAssertionAlreadyUsed
if errors.As(err, &replayAttackErr) {
return nil, ErrReplayAttackDetected{AssertionID: assertion.ID, Err: replayAttackErr}
}
return nil, fmt.Errorf("cannot prevent replay attack: %w", err)
}
}
email, fullname, samlRole, err := ExtractUserAttributes(
assertion,
config.AttributeEmail,
config.AttributeFirstname,
config.AttributeLastname,
config.AttributeRole,
)
if err != nil {
return nil, ErrCannotExtractUserAttributes{Err: err}
}
if !strings.EqualFold(email.Domain(), config.EmailDomain) {
return nil, fmt.Errorf("email domain mismatch: assertion contains email with domain %s but SAML config is for domain %s", email.Domain(), config.EmailDomain)
}
systemRole := MapSAMLRoleToSystemRole(samlRole)
samlSubject := ""
if assertion.Subject != nil && assertion.Subject.NameID != nil {
samlSubject = assertion.Subject.NameID.Value
}
return &SAMLUserInfo{
Email: email,
FullName: fullname,
Role: systemRole,
SAMLSubject: samlSubject,
OrganizationID: config.OrganizationID,
SAMLConfigID: config.ID,
}, nil
}
func (s *SAMLService) GetMetadataURL(organizationID gid.GID) string {
return fmt.Sprintf("%s/connect/saml/metadata/%s", s.baseURL, organizationID)
}
func (s *SAMLService) GenerateMetadata() ([]byte, error) {
if s.certificate == nil {
return nil, ErrSPCertificateNotConfigured{}
}
return GenerateServiceProviderMetadata(
s.GetEntityID(),
s.GetAcsURL(),
s.certificate,
)
}

View File

@@ -1,103 +0,0 @@
// Copyright (c) 2025 Probo Inc <hello@getprobo.com>.
//
// Permission to use, copy, modify, and/or distribute this software for any
// purpose with or without fee is hereby granted, provided that the above
// copyright notice and this permission notice appear in all copies.
//
// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH
// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY
// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT,
// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM
// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR
// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR
// PERFORMANCE OF THIS SOFTWARE.
package auth
import (
"context"
"errors"
"fmt"
"time"
"github.com/crewjam/saml"
"go.probo.inc/probo/pkg/coredata"
"go.probo.inc/probo/pkg/gid"
"github.com/jackc/pgx/v5/pgconn"
"go.gearno.de/kit/pg"
)
func PreventReplayAttack(
ctx context.Context,
conn pg.Conn,
scope coredata.Scoper,
assertionID string,
organizationID gid.GID,
expiresAt time.Time,
) error {
now := time.Now()
assertion := coredata.SAMLAssertion{
ID: assertionID,
OrganizationID: organizationID,
UsedAt: now,
ExpiresAt: expiresAt,
}
if err := assertion.Insert(ctx, conn, scope); err != nil {
var pgErr *pgconn.PgError
if errors.As(err, &pgErr) && pgErr.Code == "23505" && pgErr.ConstraintName == "auth_saml_assertions_pkey" {
return coredata.ErrAssertionAlreadyUsed{AssertionID: assertionID}
}
return fmt.Errorf("cannot store assertion ID: %w", err)
}
return nil
}
func ValidateAssertion(
assertion *saml.Assertion,
expectedAudience string,
now time.Time,
) error {
const clockSkewTolerance = 5 * time.Minute
if assertion.Conditions != nil && !assertion.Conditions.NotBefore.IsZero() {
if now.Add(clockSkewTolerance).Before(assertion.Conditions.NotBefore) {
return fmt.Errorf("assertion not yet valid (NotBefore: %v, now: %v, tolerance: %v)",
assertion.Conditions.NotBefore, now, clockSkewTolerance)
}
}
if assertion.Conditions != nil && !assertion.Conditions.NotOnOrAfter.IsZero() {
if now.Add(-clockSkewTolerance).After(assertion.Conditions.NotOnOrAfter) ||
now.Add(-clockSkewTolerance).Equal(assertion.Conditions.NotOnOrAfter) {
return fmt.Errorf("assertion expired (NotOnOrAfter: %v, now: %v, tolerance: %v)",
assertion.Conditions.NotOnOrAfter, now, clockSkewTolerance)
}
}
if assertion.Conditions != nil && len(assertion.Conditions.AudienceRestrictions) > 0 {
audienceValid := false
for _, restriction := range assertion.Conditions.AudienceRestrictions {
if restriction.Audience.Value == expectedAudience {
audienceValid = true
break
}
}
if !audienceValid {
return fmt.Errorf("assertion audience restriction does not match expected audience %q", expectedAudience)
}
}
return nil
}
func CleanupExpiredAssertions(ctx context.Context, conn pg.Conn) (int64, error) {
return coredata.DeleteExpiredSAMLAssertions(ctx, conn, time.Now())
}
func CleanupExpiredRequests(ctx context.Context, conn pg.Conn) (int64, error) {
return coredata.DeleteExpiredSAMLRequests(ctx, conn, time.Now())
}

File diff suppressed because it is too large Load Diff

View File

@@ -1,933 +0,0 @@
// Copyright (c) 2025 Probo Inc <hello@getprobo.com>.
//
// Permission to use, copy, modify, and/or distribute this software for any
// purpose with or without fee is hereby granted, provided that the above
// copyright notice and this permission notice appear in all copies.
//
// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH
// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY
// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT,
// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM
// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR
// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR
// PERFORMANCE OF THIS SOFTWARE.
package authz
import (
"context"
"errors"
"fmt"
"net/url"
"slices"
"time"
"go.gearno.de/kit/pg"
"go.probo.inc/probo/packages/emails"
"go.probo.inc/probo/pkg/coredata"
"go.probo.inc/probo/pkg/gid"
"go.probo.inc/probo/pkg/mail"
"go.probo.inc/probo/pkg/page"
"go.probo.inc/probo/pkg/statelesstoken"
)
type TenantAccessError struct {
Message string
}
func (e *TenantAccessError) Error() string {
return "not authorized"
}
type PermissionDeniedError struct {
Message string
}
func (e *PermissionDeniedError) Error() string {
return e.Message
}
type (
Service struct {
pg *pg.Client
baseURL string
tokenSecret string
invitationTokenValidity time.Duration
}
TenantAuthzService struct {
pg *pg.Client
baseURL string
tokenSecret string
invitationTokenValidity time.Duration
scope coredata.Scoper
}
)
const (
TokenTypeOrganizationInvitation = "organization_invitation"
)
func NewService(
ctx context.Context,
pgClient *pg.Client,
baseURL string,
tokenSecret string,
invitationTokenValidity time.Duration,
) (*Service, error) {
return &Service{
pg: pgClient,
baseURL: baseURL,
tokenSecret: tokenSecret,
invitationTokenValidity: invitationTokenValidity,
}, nil
}
func (s *Service) WithTenant(tenantID gid.TenantID) *TenantAuthzService {
return &TenantAuthzService{
pg: s.pg,
baseURL: s.baseURL,
tokenSecret: s.tokenSecret,
invitationTokenValidity: s.invitationTokenValidity,
scope: coredata.NewScope(tenantID),
}
}
func (s *Service) GetAllUserOrganizations(
ctx context.Context,
userID gid.GID,
) (coredata.Organizations, error) {
organizations := coredata.Organizations{}
err := s.pg.WithConn(
ctx,
func(conn pg.Conn) error {
if err := organizations.LoadAllByUserID(ctx, conn, userID); err != nil {
return fmt.Errorf("cannot load user organizations: %w", err)
}
return nil
},
)
return organizations, err
}
func (s *Service) GetUserOrganizationsWithRole(
ctx context.Context,
userID gid.GID,
role coredata.MembershipRole,
) (coredata.Organizations, error) {
organizations := coredata.Organizations{}
err := s.pg.WithConn(
ctx,
func(conn pg.Conn) error {
if err := organizations.LoadAllByUserIDWithRole(ctx, conn, userID, role); err != nil {
return fmt.Errorf("cannot load user organizations with role: %w", err)
}
return nil
},
)
return organizations, err
}
func (s *Service) GetAllOrganizationsForUserAPIKeyId(
ctx context.Context,
userAPIKeyID gid.GID,
) (coredata.Organizations, error) {
organizations := coredata.Organizations{}
err := s.pg.WithConn(
ctx,
func(conn pg.Conn) error {
if err := organizations.LoadAllByUserAPIKeyID(ctx, conn, userAPIKeyID); err != nil {
return fmt.Errorf("cannot load user api key organizations: %w", err)
}
return nil
},
)
return organizations, err
}
func (s *Service) GetUserOrganizations(
ctx context.Context,
userID gid.GID,
cursor *page.Cursor[coredata.OrganizationOrderField],
) (coredata.Organizations, error) {
organizations := coredata.Organizations{}
err := s.pg.WithConn(
ctx,
func(conn pg.Conn) error {
if err := organizations.LoadByUserID(ctx, conn, coredata.NewNoScope(), userID, cursor); err != nil {
return fmt.Errorf("cannot load user organizations: %w", err)
}
return nil
},
)
return organizations, err
}
func (s *Service) AcceptInvitationByID(
ctx context.Context,
invitationID gid.GID,
userID gid.GID,
) (*coredata.Invitation, error) {
var acceptedInvitation *coredata.Invitation
scope := coredata.NewScope(invitationID.TenantID())
err := s.pg.WithTx(
ctx,
func(tx pg.Conn) error {
invitation := &coredata.Invitation{}
if err := invitation.LoadByID(ctx, tx, scope, invitationID); err != nil {
var errInvitationNotFound *coredata.ErrInvitationNotFound
if errors.As(err, &errInvitationNotFound) {
return fmt.Errorf("invitation was deleted or no longer exists")
}
return fmt.Errorf("cannot load invitation: %w", err)
}
if invitation.AcceptedAt != nil {
return fmt.Errorf("invitation already accepted")
}
if time.Now().After(invitation.ExpiresAt) {
return fmt.Errorf("invitation expired")
}
user := &coredata.User{}
if err := user.LoadByID(ctx, tx, userID); err != nil {
return fmt.Errorf("cannot load user: %w", err)
}
if invitation.Email != user.EmailAddress {
return fmt.Errorf("invitation email does not match user email")
}
now := time.Now()
membershipID := gid.New(scope.GetTenantID(), coredata.MembershipEntityType)
membership := &coredata.Membership{
ID: membershipID,
UserID: userID,
OrganizationID: invitation.OrganizationID,
Role: invitation.Role,
CreatedAt: now,
UpdatedAt: now,
}
if err := membership.Create(ctx, tx, scope); err != nil {
return fmt.Errorf("cannot add user to organization: %w", err)
}
invitation.AcceptedAt = &now
if err := invitation.Update(ctx, tx, scope); err != nil {
return fmt.Errorf("cannot mark invitation as accepted: %w", err)
}
acceptedInvitation = invitation
return nil
},
)
if err != nil {
return nil, err
}
return acceptedInvitation, nil
}
type UserInvitation struct {
ID gid.GID
Email mail.Addr
FullName string
Role coredata.MembershipRole
ExpiresAt time.Time
AcceptedAt *time.Time
CreatedAt time.Time
OrganizationID gid.GID
Organization OrganizationSummary
}
type OrganizationSummary struct {
ID gid.GID
Name string
}
func (s *Service) GetUserPendingInvitations(
ctx context.Context,
email mail.Addr,
) ([]*UserInvitation, error) {
userInvitations := []*UserInvitation{}
err := s.pg.WithConn(
ctx,
func(conn pg.Conn) error {
cursor := page.NewCursor(
1000,
nil,
page.Head,
page.OrderBy[coredata.InvitationOrderField]{
Field: coredata.InvitationOrderFieldCreatedAt,
Direction: page.OrderDirectionDesc,
},
)
filter := coredata.NewInvitationFilter([]coredata.InvitationStatus{coredata.InvitationStatusPending})
invitations := coredata.Invitations{}
if err := invitations.LoadByEmail(ctx, conn, coredata.NewNoScope(), email, cursor, filter); err != nil {
return fmt.Errorf("cannot load invitations: %w", err)
}
organizationIDs := []gid.GID{}
for _, invitation := range invitations {
organizationIDs = append(organizationIDs, invitation.OrganizationID)
}
organizations := coredata.Organizations{}
if err := organizations.BatchLoadByID(ctx, conn, coredata.NewNoScope(), organizationIDs); err != nil {
return fmt.Errorf("cannot load organizations: %w", err)
}
for _, invitation := range invitations {
userInvitation := &UserInvitation{
ID: invitation.ID,
Email: invitation.Email,
FullName: invitation.FullName,
Role: invitation.Role,
ExpiresAt: invitation.ExpiresAt,
AcceptedAt: invitation.AcceptedAt,
CreatedAt: invitation.CreatedAt,
OrganizationID: invitation.OrganizationID,
}
for _, org := range organizations {
if org.ID == invitation.OrganizationID {
userInvitation.Organization = OrganizationSummary{
ID: org.ID,
Name: org.Name,
}
}
}
userInvitations = append(userInvitations, userInvitation)
}
return nil
},
)
if err != nil {
return nil, err
}
return userInvitations, nil
}
func (s *TenantAuthzService) GetOrganizationByInvitationID(
ctx context.Context,
invitationID gid.GID,
) (*coredata.Organization, error) {
var organization coredata.Organization
err := s.pg.WithConn(
ctx,
func(conn pg.Conn) error {
var invitation coredata.Invitation
if err := invitation.LoadByID(ctx, conn, s.scope, invitationID); err != nil {
return fmt.Errorf("cannot load invitation: %w", err)
}
if err := organization.LoadByID(ctx, conn, s.scope, invitation.OrganizationID); err != nil {
return fmt.Errorf("cannot load organization: %w", err)
}
return nil
},
)
if err != nil {
return nil, err
}
return &organization, nil
}
func (s *TenantAuthzService) AddUserToOrganization(
ctx context.Context,
userID gid.GID,
orgID gid.GID,
role coredata.MembershipRole,
) error {
now := time.Now()
membershipID := gid.New(s.scope.GetTenantID(), coredata.MembershipEntityType)
membership := &coredata.Membership{
ID: membershipID,
UserID: userID,
OrganizationID: orgID,
Role: role,
CreatedAt: now,
UpdatedAt: now,
}
return s.pg.WithConn(
ctx,
func(conn pg.Conn) error {
if err := membership.Create(ctx, conn, s.scope); err != nil {
return fmt.Errorf("cannot add user to organization: %w", err)
}
return nil
},
)
}
func (s *TenantAuthzService) GetInvitationsByOrganizationID(
ctx context.Context,
orgID gid.GID,
cursor *page.Cursor[coredata.InvitationOrderField],
filter *coredata.InvitationFilter,
) (*page.Page[*coredata.Invitation, coredata.InvitationOrderField], error) {
var invitations coredata.Invitations
err := s.pg.WithConn(
ctx,
func(conn pg.Conn) error {
if err := invitations.LoadByOrganizationID(ctx, conn, s.scope, orgID, cursor, filter); err != nil {
return fmt.Errorf("cannot load organization invitations: %w", err)
}
return nil
},
)
if err != nil {
return nil, err
}
return page.NewPage(invitations, cursor), nil
}
func (s *TenantAuthzService) CountOrganizationInvitations(
ctx context.Context,
orgID gid.GID,
filter *coredata.InvitationFilter,
) (int, error) {
var count int
err := s.pg.WithConn(
ctx,
func(conn pg.Conn) (err error) {
var invitations coredata.Invitations
count, err = invitations.CountByOrganizationID(ctx, conn, s.scope, orgID, filter)
if err != nil {
return fmt.Errorf("cannot count organization invitations: %w", err)
}
return nil
},
)
if err != nil {
return 0, err
}
return count, nil
}
func (s *TenantAuthzService) GetInvitationByID(
ctx context.Context,
invitationID gid.GID,
) (*coredata.Invitation, error) {
invitation := &coredata.Invitation{}
err := s.pg.WithConn(
ctx,
func(conn pg.Conn) error {
if err := invitation.LoadByID(ctx, conn, s.scope, invitationID); err != nil {
return fmt.Errorf("cannot load invitation: %w", err)
}
return nil
},
)
if err != nil {
return nil, err
}
return invitation, nil
}
func (s *TenantAuthzService) DeleteInvitation(
ctx context.Context,
invitationID gid.GID,
) error {
return s.pg.WithConn(
ctx,
func(conn pg.Conn) error {
invitation := &coredata.Invitation{}
if err := invitation.LoadByID(ctx, conn, s.scope, invitationID); err != nil {
return fmt.Errorf("cannot load invitation: %w", err)
}
if err := invitation.Delete(ctx, conn, s.scope); err != nil {
return fmt.Errorf("cannot delete invitation: %w", err)
}
return nil
},
)
}
func (s *TenantAuthzService) GetMembershipByUserAndOrganizationID(
ctx context.Context,
userID gid.GID,
orgID gid.GID,
) (*coredata.Membership, error) {
membership := &coredata.Membership{}
err := s.pg.WithConn(
ctx,
func(conn pg.Conn) error {
if err := membership.LoadByUserAndOrg(ctx, conn, s.scope, userID, orgID); err != nil {
return fmt.Errorf("cannot load membership: %w", err)
}
return nil
},
)
if err != nil {
return nil, err
}
return membership, nil
}
func (s *TenantAuthzService) GetMembershipsByOrganizationID(
ctx context.Context,
orgID gid.GID,
cursor *page.Cursor[coredata.MembershipOrderField],
) (*page.Page[*coredata.Membership, coredata.MembershipOrderField], error) {
var memberships coredata.Memberships
err := s.pg.WithConn(
ctx,
func(conn pg.Conn) error {
if err := memberships.LoadByOrganizationID(ctx, conn, s.scope, orgID, cursor); err != nil {
return fmt.Errorf("cannot load organization memberships: %w", err)
}
return nil
},
)
if err != nil {
return nil, err
}
return page.NewPage(memberships, cursor), nil
}
func (s *TenantAuthzService) CountOrganizationMemberships(
ctx context.Context,
orgID gid.GID,
) (int, error) {
var count int
err := s.pg.WithConn(
ctx,
func(conn pg.Conn) error {
var memberships coredata.Memberships
var err error
count, err = memberships.CountByOrganizationID(ctx, conn, s.scope, orgID)
return err
},
)
if err != nil {
return 0, fmt.Errorf("cannot count memberships: %w", err)
}
return count, nil
}
func (s *TenantAuthzService) CountOrganizationUsers(
ctx context.Context,
orgID gid.GID,
) (int, error) {
var count int
err := s.pg.WithConn(
ctx,
func(conn pg.Conn) (err error) {
var users coredata.Users
count, err = users.CountByOrganizationID(ctx, conn, s.scope, orgID)
if err != nil {
return fmt.Errorf("cannot count organization users: %w", err)
}
return nil
},
)
if err != nil {
return 0, fmt.Errorf("cannot count users: %w", err)
}
return count, nil
}
func (s *TenantAuthzService) GetUserRoleInOrganization(
ctx context.Context,
userID gid.GID,
orgID gid.GID,
) (coredata.MembershipRole, error) {
membership := &coredata.Membership{}
err := s.pg.WithConn(
ctx,
func(conn pg.Conn) error {
if err := membership.LoadByUserAndOrg(ctx, conn, s.scope, userID, orgID); err != nil {
return fmt.Errorf("cannot get user role: %w", err)
}
return nil
},
)
if err != nil {
return "", err
}
return membership.Role, nil
}
func (s *TenantAuthzService) RemoveMemberFromOrganization(
ctx context.Context,
orgID gid.GID,
memberID gid.GID,
) error {
membership := &coredata.Membership{}
return s.pg.WithTx(
ctx,
func(tx pg.Conn) error {
if err := membership.LoadByID(ctx, tx, s.scope, memberID); err != nil {
return fmt.Errorf("cannot load membership: %w", err)
}
if membership.OrganizationID != orgID {
return fmt.Errorf("membership does not belong to organization")
}
if err := membership.Delete(ctx, tx, s.scope); err != nil {
return fmt.Errorf("cannot delete membership: %w", err)
}
return nil
},
)
}
func (s *TenantAuthzService) UpdateMembershipRole(
ctx context.Context,
orgID gid.GID,
memberID gid.GID,
newRole coredata.MembershipRole,
) (*coredata.Membership, error) {
membership := &coredata.Membership{}
err := s.pg.WithTx(
ctx,
func(tx pg.Conn) error {
if err := membership.LoadByID(ctx, tx, s.scope, memberID); err != nil {
return fmt.Errorf("cannot load membership: %w", err)
}
if membership.OrganizationID != orgID {
return fmt.Errorf("membership does not belong to organization")
}
// If the new role cannot create API keys, delete all related API key memberships
if newRole != coredata.MembershipRoleOwner {
var apiKeyMemberships coredata.UserAPIKeyMemberships
if err := apiKeyMemberships.LoadByMembershipID(ctx, tx, s.scope, memberID); err != nil {
return fmt.Errorf("cannot load api key memberships: %w", err)
}
for _, apiKeyMembership := range apiKeyMemberships {
if err := apiKeyMembership.Delete(ctx, tx, s.scope); err != nil {
return fmt.Errorf("cannot delete api key membership: %w", err)
}
}
}
membership.Role = newRole
membership.UpdatedAt = time.Now()
if err := membership.Update(ctx, tx, s.scope); err != nil {
return fmt.Errorf("cannot update membership role: %w", err)
}
return nil
},
)
if err != nil {
return nil, err
}
return membership, nil
}
func (s *TenantAuthzService) InviteUserToOrganization(
ctx context.Context,
organizationID gid.GID,
emailAddress mail.Addr,
fullName string,
role coredata.MembershipRole,
) (*coredata.Invitation, error) {
var invitation *coredata.Invitation
err := s.pg.WithTx(
ctx,
func(tx pg.Conn) error {
user := &coredata.User{}
userExists := true
if err := user.LoadByEmail(ctx, tx, emailAddress); err != nil {
var userNotFound *coredata.ErrUserNotFound
if errors.As(err, &userNotFound) {
userExists = false
} else {
return fmt.Errorf("cannot check if user exists: %w", err)
}
}
organization := &coredata.Organization{}
if err := organization.LoadByID(ctx, tx, s.scope, organizationID); err != nil {
return fmt.Errorf("cannot load organization: %w", err)
}
invitationID := gid.New(s.scope.GetTenantID(), coredata.InvitationEntityType)
now := time.Now()
invitation = &coredata.Invitation{
ID: invitationID,
OrganizationID: organizationID,
Email: emailAddress,
FullName: fullName,
Role: role,
ExpiresAt: now.Add(s.invitationTokenValidity),
CreatedAt: now,
}
var err error
var invitationURL string
var recipientName string
if userExists {
recipientName = user.FullName
invitationURL = s.baseURL + "/"
} else {
recipientName = fullName
invitationData := coredata.InvitationData{
InvitationID: invitationID,
OrganizationID: organizationID,
Email: emailAddress,
FullName: fullName,
Role: role,
}
invitationToken, err := statelesstoken.NewToken(
s.tokenSecret,
TokenTypeOrganizationInvitation,
s.invitationTokenValidity,
invitationData,
)
if err != nil {
return fmt.Errorf("cannot generate invitation token: %w", err)
}
invitationURL = fmt.Sprintf("%s/auth/signup-from-invitation?token=%s&fullName=%s", s.baseURL, invitationToken, url.QueryEscape(fullName))
}
subject, textBody, htmlBody, err := emails.RenderInvitation(
s.baseURL,
recipientName,
organization.Name,
invitationURL,
)
if err != nil {
return fmt.Errorf("cannot render invitation email: %w", err)
}
email := coredata.NewEmail(
fullName,
emailAddress,
subject,
textBody,
htmlBody,
)
if err := email.Insert(ctx, tx); err != nil {
return fmt.Errorf("cannot insert email: %w", err)
}
if err := invitation.Create(ctx, tx, s.scope); err != nil {
return fmt.Errorf("cannot create invitation: %w", err)
}
return nil
},
)
if err != nil {
return nil, err
}
return invitation, nil
}
func (s *TenantAuthzService) EnsureSAMLMembership(
ctx context.Context,
userID gid.GID,
organizationID gid.GID,
role *coredata.MembershipRole,
) error {
now := time.Now()
return s.pg.WithTx(
ctx,
func(tx pg.Conn) error {
var membership coredata.Membership
err := membership.LoadByUserAndOrg(ctx, tx, s.scope, userID, organizationID)
if err != nil {
if _, ok := err.(coredata.ErrMembershipNotFound); !ok {
return fmt.Errorf("cannot load membership: %w", err)
}
membershipRole := coredata.MembershipRoleViewer
if role != nil {
membershipRole = *role
}
membershipID := gid.New(s.scope.GetTenantID(), coredata.MembershipEntityType)
membership = coredata.Membership{
ID: membershipID,
UserID: userID,
OrganizationID: organizationID,
Role: membershipRole,
CreatedAt: now,
UpdatedAt: now,
}
if err := membership.Create(ctx, tx, s.scope); err != nil {
return fmt.Errorf("cannot create membership: %w", err)
}
return nil
}
if role != nil && membership.Role != *role {
membership.Role = *role
membership.UpdatedAt = now
if err := membership.Update(ctx, tx, s.scope); err != nil {
return fmt.Errorf("cannot update membership role: %w", err)
}
}
return nil
},
)
}
func (s *TenantAuthzService) Authorize(
ctx context.Context,
user *coredata.User,
apiKey *coredata.UserAPIKey,
entityGID gid.GID,
action Action,
) error {
requiredRoles := GetPermissionsForAction(entityGID.EntityType(), action)
if requiredRoles == nil {
entityModel, _ := coredata.EntityModel(entityGID.EntityType())
return &PermissionDeniedError{
Message: fmt.Sprintf("no permissions defined for action %s on entity %s", action, entityModel),
}
}
role, err := s.GetUserOrAPIKeyRole(ctx, user, apiKey, entityGID)
if err != nil {
return fmt.Errorf("cannot get user or API key role: %w", err)
}
if !slices.Contains(requiredRoles, role) {
return &PermissionDeniedError{
Message: fmt.Sprintf("role %s not authorized for action %s, requires one of %v", role, action, requiredRoles),
}
}
return nil
}
func (s *TenantAuthzService) CanAssignRole(
ctx context.Context,
user *coredata.User,
apiKey *coredata.UserAPIKey,
entityGID gid.GID,
targetRole coredata.MembershipRole,
) error {
currentRole, err := s.GetUserOrAPIKeyRole(ctx, user, apiKey, entityGID)
if err != nil {
return fmt.Errorf("cannot get user or API key role: %w", err)
}
if currentRole == RoleOwner || currentRole == RoleFull {
return nil
}
if currentRole == RoleAdmin {
if targetRole == coredata.MembershipRoleOwner {
return &PermissionDeniedError{Message: "admin users cannot assign owner role"}
}
return nil
}
return &PermissionDeniedError{Message: fmt.Sprintf("role %s cannot assign roles", currentRole)}
}
func (s *TenantAuthzService) GetUserOrAPIKeyRole(
ctx context.Context,
user *coredata.User,
apiKey *coredata.UserAPIKey,
entityGID gid.GID,
) (Role, error) {
var role Role
err := s.pg.WithConn(
ctx,
func(conn pg.Conn) error {
if user != nil {
membership := &coredata.Membership{}
if err := membership.LoadRoleByUserAndEntityID(ctx, conn, s.scope, user.ID, entityGID); err != nil {
return fmt.Errorf("cannot get user role: %w", err)
}
role = Role(membership.Role.String())
return nil
}
if apiKey != nil {
apiKeyMembership := &coredata.UserAPIKeyMembership{}
if err := apiKeyMembership.LoadRoleByAPIKeyAndEntityID(ctx, conn, s.scope, apiKey.ID, entityGID); err != nil {
return fmt.Errorf("cannot get API key role: %w", err)
}
role = Role(apiKeyMembership.Role.String())
return nil
}
return fmt.Errorf("no user or API key provided")
},
)
if err != nil {
return "", err
}
return role, nil
}

View File

@@ -109,10 +109,10 @@ func (b *BaseURL) Port() string {
// URLBuilder provides a fluent interface for building URLs. // URLBuilder provides a fluent interface for building URLs.
type URLBuilder struct { type URLBuilder struct {
base *BaseURL base *BaseURL
path string path string
query url.Values query url.Values
err error err error
} }
// WithPath returns a URLBuilder with the specified path. // WithPath returns a URLBuilder with the specified path.
@@ -224,3 +224,12 @@ func (b *BaseURL) MarshalText() ([]byte, error) {
} }
return []byte(b.raw), nil return []byte(b.raw), nil
} }
func (b *URLBuilder) URL() url.URL {
return url.URL{
Scheme: b.base.Scheme(),
Host: b.base.Host(),
Path: b.path,
RawQuery: b.query.Encode(),
}
}

View File

@@ -0,0 +1,120 @@
// Copyright (c) 2025 Probo Inc <hello@getprobo.com>.
//
// Permission to use, copy, modify, and/or distribute this software for any
// purpose with or without fee is hereby granted, provided that the above
// copyright notice and this permission notice appear in all copies.
//
// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH
// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY
// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT,
// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM
// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR
// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR
// PERFORMANCE OF THIS SOFTWARE.
// Package bearertoken parses Bearer tokens according to RFC 6750.
//
// The grammar is defined as:
//
// b64token = 1*( ALPHA / DIGIT / "-" / "." / "_" / "~" / "+" / "/" ) *"="
// credentials = "Bearer" 1*SP b64token
package bearertoken
import (
"errors"
"strings"
)
var (
// ErrInvalidCredentials is returned when the credentials string is malformed.
ErrInvalidCredentials = errors.New("invalid bearer credentials")
// ErrMissingToken is returned when the token part is empty.
ErrMissingToken = errors.New("missing bearer token")
// ErrInvalidToken is returned when the token contains invalid characters.
ErrInvalidToken = errors.New("invalid bearer token")
)
const (
scheme = "Bearer"
)
// Parse extracts the b64token from a Bearer credentials string.
// The input must follow the format: "Bearer" 1*SP b64token
func Parse(credentials string) (string, error) {
if len(credentials) <= len(scheme) {
return "", ErrInvalidCredentials
}
if !strings.EqualFold(credentials[:len(scheme)], scheme) {
return "", ErrInvalidCredentials
}
rest := credentials[len(scheme):]
if len(rest) == 0 || rest[0] != ' ' {
return "", ErrInvalidCredentials
}
// Skip all spaces (1*SP)
token := strings.TrimLeft(rest, " ")
if token == "" {
return "", ErrMissingToken
}
if !isValidToken(token) {
return "", ErrInvalidToken
}
return token, nil
}
// isValidToken checks if the given string is a valid b64token.
// A valid b64token consists of 1 or more characters from the set
// [A-Za-z0-9-._~+/] followed by zero or more '=' characters.
func isValidToken(token string) bool {
if len(token) == 0 {
return false
}
// Find where the padding starts (if any)
paddingStart := strings.IndexByte(token, '=')
if paddingStart == -1 {
paddingStart = len(token)
}
// Must have at least one non-padding character
if paddingStart == 0 {
return false
}
// Validate the base part (before padding)
for i := 0; i < paddingStart; i++ {
if !isB64Char(token[i]) {
return false
}
}
// Validate padding (only '=' allowed after first '=')
for i := paddingStart; i < len(token); i++ {
if token[i] != '=' {
return false
}
}
return true
}
// isB64Char returns true if c is a valid b64token character (excluding padding).
// Valid characters: ALPHA / DIGIT / "-" / "." / "_" / "~" / "+" / "/"
func isB64Char(c byte) bool {
return (c >= 'A' && c <= 'Z') ||
(c >= 'a' && c <= 'z') ||
(c >= '0' && c <= '9') ||
c == '-' ||
c == '.' ||
c == '_' ||
c == '~' ||
c == '+' ||
c == '/'
}

View File

@@ -0,0 +1,298 @@
// Copyright (c) 2025 Probo Inc <hello@getprobo.com>.
//
// Permission to use, copy, modify, and/or distribute this software for any
// purpose with or without fee is hereby granted, provided that the above
// copyright notice and this permission notice appear in all copies.
//
// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH
// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY
// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT,
// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM
// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR
// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR
// PERFORMANCE OF THIS SOFTWARE.
package bearertoken
import (
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func TestParse(t *testing.T) {
t.Parallel()
tests := []struct {
name string
credentials string
wantToken string
wantErr error
}{
// Valid credentials
{
name: "valid simple token",
credentials: "Bearer abc123",
wantToken: "abc123",
wantErr: nil,
},
{
name: "valid uppercase token",
credentials: "Bearer ABCXYZ",
wantToken: "ABCXYZ",
wantErr: nil,
},
{
name: "valid digits only token",
credentials: "Bearer 0123456789",
wantToken: "0123456789",
wantErr: nil,
},
{
name: "valid base64 token with single padding",
credentials: "Bearer dXNlcm5hbWU6cGFzc3dvcmQ=",
wantToken: "dXNlcm5hbWU6cGFzc3dvcmQ=",
wantErr: nil,
},
{
name: "valid base64 token with double padding",
credentials: "Bearer YWJj==",
wantToken: "YWJj==",
wantErr: nil,
},
{
name: "valid token with all special chars",
credentials: "Bearer abc-._~+/123",
wantToken: "abc-._~+/123",
wantErr: nil,
},
{
name: "valid token with hyphen",
credentials: "Bearer abc-def",
wantToken: "abc-def",
wantErr: nil,
},
{
name: "valid token with dot",
credentials: "Bearer abc.def",
wantToken: "abc.def",
wantErr: nil,
},
{
name: "valid token with underscore",
credentials: "Bearer abc_def",
wantToken: "abc_def",
wantErr: nil,
},
{
name: "valid token with tilde",
credentials: "Bearer abc~def",
wantToken: "abc~def",
wantErr: nil,
},
{
name: "valid token with plus",
credentials: "Bearer abc+def",
wantToken: "abc+def",
wantErr: nil,
},
{
name: "valid token with slash",
credentials: "Bearer abc/def",
wantToken: "abc/def",
wantErr: nil,
},
{
name: "valid token with multiple spaces after scheme",
credentials: "Bearer token123",
wantToken: "token123",
wantErr: nil,
},
{
name: "valid jwt-like token",
credentials: "Bearer eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJzdWIiOiIxMjM0NTY3ODkwIn0.dozjgNryP4J3jVmNHl0w5N_XgL0n3I9PlFUP0THsR8U",
wantToken: "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJzdWIiOiIxMjM0NTY3ODkwIn0.dozjgNryP4J3jVmNHl0w5N_XgL0n3I9PlFUP0THsR8U",
wantErr: nil,
},
{
name: "valid single char token",
credentials: "Bearer a",
wantToken: "a",
wantErr: nil,
},
// Case insensitive scheme
{
name: "lowercase scheme",
credentials: "bearer abc123",
wantToken: "abc123",
wantErr: nil,
},
{
name: "uppercase scheme",
credentials: "BEARER abc123",
wantToken: "abc123",
wantErr: nil,
},
{
name: "mixed case scheme",
credentials: "BeArEr abc123",
wantToken: "abc123",
wantErr: nil,
},
// Invalid credentials (scheme errors)
{
name: "empty string",
credentials: "",
wantToken: "",
wantErr: ErrInvalidCredentials,
},
{
name: "only scheme without space",
credentials: "Bearer",
wantToken: "",
wantErr: ErrInvalidCredentials,
},
{
name: "scheme without space before token",
credentials: "Bearerabc123",
wantToken: "",
wantErr: ErrInvalidCredentials,
},
{
name: "wrong scheme Basic",
credentials: "Basic abc123",
wantToken: "",
wantErr: ErrInvalidCredentials,
},
{
name: "wrong scheme Digest",
credentials: "Digest abc123",
wantToken: "",
wantErr: ErrInvalidCredentials,
},
{
name: "partial scheme",
credentials: "Bear abc123",
wantToken: "",
wantErr: ErrInvalidCredentials,
},
{
name: "scheme with tab instead of space",
credentials: "Bearer\tabc123",
wantToken: "",
wantErr: ErrInvalidCredentials,
},
// Missing token
{
name: "missing token after single space",
credentials: "Bearer ",
wantToken: "",
wantErr: ErrMissingToken,
},
{
name: "missing token after multiple spaces",
credentials: "Bearer ",
wantToken: "",
wantErr: ErrMissingToken,
},
// Invalid token characters
{
name: "invalid char @",
credentials: "Bearer abc@123",
wantToken: "",
wantErr: ErrInvalidToken,
},
{
name: "invalid char #",
credentials: "Bearer abc#123",
wantToken: "",
wantErr: ErrInvalidToken,
},
{
name: "invalid char !",
credentials: "Bearer abc!123",
wantToken: "",
wantErr: ErrInvalidToken,
},
{
name: "invalid char ?",
credentials: "Bearer abc?123",
wantToken: "",
wantErr: ErrInvalidToken,
},
{
name: "invalid char *",
credentials: "Bearer abc*123",
wantToken: "",
wantErr: ErrInvalidToken,
},
{
name: "invalid char space in token",
credentials: "Bearer abc 123",
wantToken: "",
wantErr: ErrInvalidToken,
},
{
name: "invalid char tab in token",
credentials: "Bearer abc\t123",
wantToken: "",
wantErr: ErrInvalidToken,
},
{
name: "invalid char newline in token",
credentials: "Bearer abc\n123",
wantToken: "",
wantErr: ErrInvalidToken,
},
// Invalid padding
{
name: "token starting with equals",
credentials: "Bearer =abc",
wantToken: "",
wantErr: ErrInvalidToken,
},
{
name: "token with equals in middle",
credentials: "Bearer abc=def",
wantToken: "",
wantErr: ErrInvalidToken,
},
{
name: "token only padding",
credentials: "Bearer ==",
wantToken: "",
wantErr: ErrInvalidToken,
},
{
name: "token with char after padding",
credentials: "Bearer abc==def",
wantToken: "",
wantErr: ErrInvalidToken,
},
}
for _, tt := range tests {
t.Run(
tt.name,
func(t *testing.T) {
t.Parallel()
gotToken, gotErr := Parse(tt.credentials)
if tt.wantErr != nil {
require.ErrorIs(t, gotErr, tt.wantErr)
assert.Empty(t, gotToken)
} else {
require.NoError(t, gotErr)
assert.Equal(t, tt.wantToken, gotToken)
}
},
)
}
}

View File

@@ -21,10 +21,10 @@ import (
"maps" "maps"
"time" "time"
"go.probo.inc/probo/pkg/gid"
"go.probo.inc/probo/pkg/page"
"github.com/jackc/pgx/v5" "github.com/jackc/pgx/v5"
"go.gearno.de/kit/pg" "go.gearno.de/kit/pg"
"go.probo.inc/probo/pkg/gid"
"go.probo.inc/probo/pkg/page"
) )
type ( type (
@@ -43,24 +43,8 @@ type (
} }
Assets []*Asset Assets []*Asset
ErrAssetNotFound struct {
Identifier string
}
ErrAssetAlreadyExists struct {
message string
}
) )
func (e ErrAssetNotFound) Error() string {
return fmt.Sprintf("asset not found: %q", e.Identifier)
}
func (e ErrAssetAlreadyExists) Error() string {
return e.message
}
func (a *Asset) CursorKey(field AssetOrderField) page.CursorKey { func (a *Asset) CursorKey(field AssetOrderField) page.CursorKey {
switch field { switch field {
case AssetOrderFieldCreatedAt: case AssetOrderFieldCreatedAt:
@@ -112,7 +96,7 @@ LIMIT 1;
asset, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[Asset]) asset, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[Asset])
if err != nil { if err != nil {
if errors.Is(err, pgx.ErrNoRows) { if errors.Is(err, pgx.ErrNoRows) {
return &ErrAssetNotFound{Identifier: assetID.String()} return ErrResourceNotFound
} }
return fmt.Errorf("cannot collect asset: %w", err) return fmt.Errorf("cannot collect asset: %w", err)
@@ -162,7 +146,7 @@ LIMIT 1;
asset, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[Asset]) asset, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[Asset])
if err != nil { if err != nil {
if errors.Is(err, pgx.ErrNoRows) { if errors.Is(err, pgx.ErrNoRows) {
return &ErrAssetNotFound{Identifier: a.OwnerID.String()} return ErrResourceNotFound
} }
return fmt.Errorf("cannot collect asset: %w", err) return fmt.Errorf("cannot collect asset: %w", err)

View File

@@ -21,10 +21,10 @@ import (
"maps" "maps"
"time" "time"
"go.probo.inc/probo/pkg/gid"
"go.probo.inc/probo/pkg/page"
"github.com/jackc/pgx/v5" "github.com/jackc/pgx/v5"
"go.gearno.de/kit/pg" "go.gearno.de/kit/pg"
"go.probo.inc/probo/pkg/gid"
"go.probo.inc/probo/pkg/page"
) )
type ( type (
@@ -43,24 +43,8 @@ type (
} }
Audits []*Audit Audits []*Audit
ErrAuditNotFound struct {
Identifier string
}
ErrAuditAlreadyExists struct {
message string
}
) )
func (e ErrAuditNotFound) Error() string {
return fmt.Sprintf("audit not found: %q", e.Identifier)
}
func (e ErrAuditAlreadyExists) Error() string {
return e.message
}
func (a *Audit) CursorKey(field AuditOrderField) page.CursorKey { func (a *Audit) CursorKey(field AuditOrderField) page.CursorKey {
switch field { switch field {
case AuditOrderFieldCreatedAt: case AuditOrderFieldCreatedAt:
@@ -116,7 +100,7 @@ LIMIT 1;
audit, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[Audit]) audit, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[Audit])
if err != nil { if err != nil {
if errors.Is(err, pgx.ErrNoRows) { if errors.Is(err, pgx.ErrNoRows) {
return &ErrAuditNotFound{Identifier: auditID.String()} return ErrResourceNotFound
} }
return fmt.Errorf("cannot collect audit: %w", err) return fmt.Errorf("cannot collect audit: %w", err)
@@ -530,7 +514,7 @@ LIMIT 1;
audit, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[Audit]) audit, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[Audit])
if err != nil { if err != nil {
if errors.Is(err, pgx.ErrNoRows) { if errors.Is(err, pgx.ErrNoRows) {
return &ErrAuditNotFound{Identifier: reportID.String()} return ErrResourceNotFound
} }
return fmt.Errorf("cannot collect audit: %w", err) return fmt.Errorf("cannot collect audit: %w", err)

View File

@@ -44,24 +44,8 @@ type (
} }
Controls []*Control Controls []*Control
ErrControlNotFound struct {
Identifier string
}
ErrControlAlreadyExists struct {
message string
}
) )
func (e ErrControlNotFound) Error() string {
return fmt.Sprintf("control not found: %q", e.Identifier)
}
func (e ErrControlAlreadyExists) Error() string {
return e.message
}
func (c Control) CursorKey(orderBy ControlOrderField) page.CursorKey { func (c Control) CursorKey(orderBy ControlOrderField) page.CursorKey {
switch orderBy { switch orderBy {
case ControlOrderFieldCreatedAt: case ControlOrderFieldCreatedAt:
@@ -661,7 +645,7 @@ LIMIT 1;
control, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[Control]) control, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[Control])
if err != nil { if err != nil {
if errors.Is(err, pgx.ErrNoRows) { if errors.Is(err, pgx.ErrNoRows) {
return &ErrControlNotFound{Identifier: fmt.Sprintf("%s:%s", frameworkID, sectionTitle)} return ErrResourceNotFound
} }
return fmt.Errorf("cannot collect control: %w", err) return fmt.Errorf("cannot collect control: %w", err)
@@ -710,7 +694,7 @@ LIMIT 1;
control, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[Control]) control, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[Control])
if err != nil { if err != nil {
if errors.Is(err, pgx.ErrNoRows) { if errors.Is(err, pgx.ErrNoRows) {
return &ErrControlNotFound{Identifier: controlID.String()} return ErrResourceNotFound
} }
return fmt.Errorf("cannot collect control: %w", err) return fmt.Errorf("cannot collect control: %w", err)
@@ -778,9 +762,7 @@ VALUES (
var pgErr *pgconn.PgError var pgErr *pgconn.PgError
if errors.As(err, &pgErr) { if errors.As(err, &pgErr) {
if pgErr.Code == "23505" && pgErr.ConstraintName == "controls_framework_ref_unique" { if pgErr.Code == "23505" && pgErr.ConstraintName == "controls_framework_ref_unique" {
return &ErrControlAlreadyExists{ return ErrResourceAlreadyExists
message: fmt.Sprintf("control with framework_id %s and section_title %q already exists", c.FrameworkID, c.SectionTitle),
}
} }
} }
return fmt.Errorf("cannot insert control: %w", err) return fmt.Errorf("cannot insert control: %w", err)
@@ -848,9 +830,7 @@ WHERE %s
var pgErr *pgconn.PgError var pgErr *pgconn.PgError
if errors.As(err, &pgErr) { if errors.As(err, &pgErr) {
if pgErr.Code == "23505" && pgErr.ConstraintName == "controls_framework_ref_unique" { if pgErr.Code == "23505" && pgErr.ConstraintName == "controls_framework_ref_unique" {
return &ErrControlAlreadyExists{ return ErrResourceAlreadyExists
message: fmt.Sprintf("control with section_title %q already exists", c.SectionTitle),
}
} }
} }
return fmt.Errorf("cannot update control: %w", err) return fmt.Errorf("cannot update control: %w", err)

View File

@@ -21,10 +21,10 @@ import (
"maps" "maps"
"time" "time"
"go.probo.inc/probo/pkg/gid"
"github.com/jackc/pgx/v5" "github.com/jackc/pgx/v5"
"github.com/jackc/pgx/v5/pgconn" "github.com/jackc/pgx/v5/pgconn"
"go.gearno.de/kit/pg" "go.gearno.de/kit/pg"
"go.probo.inc/probo/pkg/gid"
) )
type ( type (
@@ -37,17 +37,8 @@ type (
} }
ControlDocuments []*ControlDocument ControlDocuments []*ControlDocument
ErrControlDocumentMappingAlreadyExists struct {
ControlID gid.GID
DocumentID gid.GID
}
) )
func (e ErrControlDocumentMappingAlreadyExists) Error() string {
return fmt.Sprintf("control %s is already mapped to document %s", e.ControlID, e.DocumentID)
}
func (cp ControlDocument) Insert( func (cp ControlDocument) Insert(
ctx context.Context, ctx context.Context,
conn pg.Conn, conn pg.Conn,
@@ -84,10 +75,7 @@ VALUES (
var pgErr *pgconn.PgError var pgErr *pgconn.PgError
if errors.As(err, &pgErr) { if errors.As(err, &pgErr) {
if pgErr.Code == "23505" && pgErr.ConstraintName == "controls_policies_pkey" { if pgErr.Code == "23505" && pgErr.ConstraintName == "controls_policies_pkey" {
return &ErrControlDocumentMappingAlreadyExists{ return ErrResourceAlreadyExists
ControlID: cp.ControlID,
DocumentID: cp.DocumentID,
}
} }
} }

View File

@@ -52,24 +52,8 @@ type (
} }
CustomDomains []*CustomDomain CustomDomains []*CustomDomain
ErrCustomDomainNotFound struct {
Identifier string
}
ErrCustomDomainAlreadyExists struct {
message string
}
) )
func (e ErrCustomDomainNotFound) Error() string {
return fmt.Sprintf("custom domain not found: %q", e.Identifier)
}
func (e ErrCustomDomainAlreadyExists) Error() string {
return e.message
}
func NewCustomDomain(tenantID gid.TenantID, domain string) *CustomDomain { func NewCustomDomain(tenantID gid.TenantID, domain string) *CustomDomain {
now := time.Now() now := time.Now()
return &CustomDomain{ return &CustomDomain{
@@ -385,9 +369,7 @@ INSERT INTO custom_domains (
var pgErr *pgconn.PgError var pgErr *pgconn.PgError
if errors.As(err, &pgErr) { if errors.As(err, &pgErr) {
if pgErr.Code == "23505" && pgErr.ConstraintName == "custom_domains_domain_key" { if pgErr.Code == "23505" && pgErr.ConstraintName == "custom_domains_domain_key" {
return &ErrCustomDomainAlreadyExists{ return ErrResourceAlreadyExists
message: fmt.Sprintf("custom domain with domain %q already exists", cd.Domain),
}
} }
} }
return fmt.Errorf("cannot insert custom domain: %w", err) return fmt.Errorf("cannot insert custom domain: %w", err)

View File

@@ -43,24 +43,8 @@ type (
} }
Documents []*Document Documents []*Document
ErrDocumentNotFound struct {
Identifier string
}
ErrDocumentAlreadyExists struct {
message string
}
) )
func (e ErrDocumentNotFound) Error() string {
return fmt.Sprintf("document not found: %q", e.Identifier)
}
func (e ErrDocumentAlreadyExists) Error() string {
return e.message
}
func (p Document) CursorKey(orderBy DocumentOrderField) page.CursorKey { func (p Document) CursorKey(orderBy DocumentOrderField) page.CursorKey {
switch orderBy { switch orderBy {
case DocumentOrderFieldCreatedAt: case DocumentOrderFieldCreatedAt:
@@ -114,7 +98,7 @@ LIMIT 1;
document, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[Document]) document, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[Document])
if err != nil { if err != nil {
if errors.Is(err, pgx.ErrNoRows) { if errors.Is(err, pgx.ErrNoRows) {
return &ErrDocumentNotFound{Identifier: documentID.String()} return ErrResourceNotFound
} }
return fmt.Errorf("cannot collect document: %w", err) return fmt.Errorf("cannot collect document: %w", err)
@@ -168,7 +152,7 @@ LIMIT 1;
document, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[Document]) document, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[Document])
if err != nil { if err != nil {
if errors.Is(err, pgx.ErrNoRows) { if errors.Is(err, pgx.ErrNoRows) {
return &ErrDocumentNotFound{Identifier: documentID.String()} return ErrResourceNotFound
} }
return fmt.Errorf("cannot collect document: %w", err) return fmt.Errorf("cannot collect document: %w", err)

View File

@@ -46,32 +46,8 @@ type (
} }
DocumentVersions []*DocumentVersion DocumentVersions []*DocumentVersion
ErrDocumentVersionNotFound struct {
Identifier string
}
ErrDocumentVersionAlreadyExists struct {
message string
}
ErrDocumentVersionNoChanges struct {
Message string
}
) )
func (e ErrDocumentVersionNotFound) Error() string {
return fmt.Sprintf("document version not found: %q", e.Identifier)
}
func (e ErrDocumentVersionAlreadyExists) Error() string {
return e.message
}
func (e ErrDocumentVersionNoChanges) Error() string {
return e.Message
}
func (p *DocumentVersions) LoadByDocumentID( func (p *DocumentVersions) LoadByDocumentID(
ctx context.Context, ctx context.Context,
conn pg.Conn, conn pg.Conn,
@@ -245,15 +221,8 @@ VALUES (
var pgErr *pgconn.PgError var pgErr *pgconn.PgError
if errors.As(err, &pgErr) { if errors.As(err, &pgErr) {
if pgErr.Code == "23505" { if pgErr.Code == "23505" {
if pgErr.ConstraintName == "document_versions_document_id_version_number_key" { if pgErr.ConstraintName == "document_versions_document_id_version_number_key" || pgErr.ConstraintName == "document_one_draft_version_idx" {
return &ErrDocumentVersionAlreadyExists{ return ErrResourceAlreadyExists
message: fmt.Sprintf("document version with document_id %s and version_number %d already exists", p.DocumentID, p.VersionNumber),
}
}
if pgErr.ConstraintName == "document_one_draft_version_idx" {
return &ErrDocumentVersionAlreadyExists{
message: fmt.Sprintf("document %s already has a draft version", p.DocumentID),
}
} }
} }
} }

View File

@@ -50,30 +50,7 @@ type (
} }
DocumentVersionSignaturesWithPeople []*DocumentVersionSignatureWithPeople DocumentVersionSignaturesWithPeople []*DocumentVersionSignatureWithPeople
ErrDocumentVersionSignatureNotFound struct {
Identifier string
}
ErrDocumentVersionSignatureAlreadyExists struct {
message string
}
ErrDocumentVersionSignatureAlreadySigned struct{}
) )
func (e ErrDocumentVersionSignatureNotFound) Error() string {
return fmt.Sprintf("document version signature not found: %q", e.Identifier)
}
func (e ErrDocumentVersionSignatureAlreadyExists) Error() string {
return e.message
}
func (e ErrDocumentVersionSignatureAlreadySigned) Error() string {
return "document version already signed"
}
func (pvs DocumentVersionSignature) CursorKey(orderBy DocumentVersionSignatureOrderField) page.CursorKey { func (pvs DocumentVersionSignature) CursorKey(orderBy DocumentVersionSignatureOrderField) page.CursorKey {
switch orderBy { switch orderBy {
case DocumentVersionSignatureOrderFieldCreatedAt: case DocumentVersionSignatureOrderFieldCreatedAt:
@@ -225,9 +202,7 @@ INSERT INTO document_version_signatures (
var pgErr *pgconn.PgError var pgErr *pgconn.PgError
if errors.As(err, &pgErr) { if errors.As(err, &pgErr) {
if pgErr.Code == "23505" && pgErr.ConstraintName == "policy_version_signatures_policy_version_id_signed_by_key" { if pgErr.Code == "23505" && pgErr.ConstraintName == "policy_version_signatures_policy_version_id_signed_by_key" {
return &ErrDocumentVersionSignatureAlreadyExists{ return ErrResourceAlreadyExists
message: fmt.Sprintf("document version signature with document_version_id %s and signed_by %s already exists", pvs.DocumentVersionID, pvs.SignedBy),
}
} }
} }
return fmt.Errorf("cannot insert document version signature: %w", err) return fmt.Errorf("cannot insert document version signature: %w", err)

11
pkg/coredata/errors.go Normal file
View File

@@ -0,0 +1,11 @@
package coredata
import (
"errors"
)
var (
ErrResourceNotFound = errors.New("resource not found")
ErrResourceAlreadyExists = errors.New("resource already exists")
ErrResourceInUse = errors.New("resource is in use")
)

View File

@@ -45,24 +45,8 @@ type (
} }
Evidences []*Evidence Evidences []*Evidence
ErrEvidenceNotFound struct {
Identifier string
}
ErrEvidenceAlreadyExists struct {
message string
}
) )
func (e ErrEvidenceNotFound) Error() string {
return fmt.Sprintf("evidence not found: %q", e.Identifier)
}
func (e ErrEvidenceAlreadyExists) Error() string {
return e.message
}
func (e Evidence) CursorKey(orderBy EvidenceOrderField) page.CursorKey { func (e Evidence) CursorKey(orderBy EvidenceOrderField) page.CursorKey {
switch orderBy { switch orderBy {
case EvidenceOrderFieldCreatedAt: case EvidenceOrderFieldCreatedAt:
@@ -192,9 +176,7 @@ VALUES (
var pgErr *pgconn.PgError var pgErr *pgconn.PgError
if errors.As(err, &pgErr) { if errors.As(err, &pgErr) {
if pgErr.Code == "23505" && pgErr.ConstraintName == "evidences_reference_id_key" { if pgErr.Code == "23505" && pgErr.ConstraintName == "evidences_reference_id_key" {
return &ErrEvidenceAlreadyExists{ return ErrResourceAlreadyExists
message: fmt.Sprintf("evidence with task_id %s and reference_id %q already exists", e.TaskID, e.ReferenceID),
}
} }
} }
return fmt.Errorf("cannot insert evidence: %w", err) return fmt.Errorf("cannot insert evidence: %w", err)

View File

@@ -0,0 +1,11 @@
package coredata
type (
ExpireReason string
)
const (
ExpireReasonIdleTimeout ExpireReason = "idle_timeout"
ExpireReasonRevoked ExpireReason = "revoked"
ExpireReasonClosed ExpireReason = "closed"
)

View File

@@ -42,24 +42,8 @@ type (
} }
Files []*File Files []*File
ErrFileNotFound struct {
Identifier string
}
ErrFileAlreadyExists struct {
message string
}
) )
func (e ErrFileNotFound) Error() string {
return fmt.Sprintf("file not found: %q", e.Identifier)
}
func (e ErrFileAlreadyExists) Error() string {
return e.message
}
func (f *File) LoadByID( func (f *File) LoadByID(
ctx context.Context, ctx context.Context,
conn pg.Conn, conn pg.Conn,
@@ -100,7 +84,7 @@ LIMIT 1;
file, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[File]) file, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[File])
if err != nil { if err != nil {
if errors.Is(err, pgx.ErrNoRows) { if errors.Is(err, pgx.ErrNoRows) {
return &ErrFileNotFound{Identifier: fileID.String()} return ErrResourceNotFound
} }
return fmt.Errorf("cannot collect file: %w", err) return fmt.Errorf("cannot collect file: %w", err)
@@ -165,9 +149,7 @@ VALUES (
var pgErr *pgconn.PgError var pgErr *pgconn.PgError
if errors.As(err, &pgErr) { if errors.As(err, &pgErr) {
if pgErr.Code == "23505" && pgErr.ConstraintName == "files_file_key_key" { if pgErr.Code == "23505" && pgErr.ConstraintName == "files_file_key_key" {
return &ErrFileAlreadyExists{ return ErrResourceAlreadyExists
message: fmt.Sprintf("file with file_key %q already exists", f.FileKey),
}
} }
} }
return fmt.Errorf("cannot insert file: %w", err) return fmt.Errorf("cannot insert file: %w", err)

View File

@@ -42,33 +42,8 @@ type (
} }
Frameworks []*Framework Frameworks []*Framework
ErrFrameworkNotFound struct {
Identifier string
}
ErrFrameworkAlreadyExists struct {
message string
}
ErrFrameworkReferenceIDAlreadyExists struct {
ReferenceID string
OrganizationID gid.GID
}
) )
func (e ErrFrameworkNotFound) Error() string {
return fmt.Sprintf("framework not found: %q", e.Identifier)
}
func (e ErrFrameworkAlreadyExists) Error() string {
return e.message
}
func (e ErrFrameworkReferenceIDAlreadyExists) Error() string {
return fmt.Sprintf("framework with reference ID %q already exists for organization %s", e.ReferenceID, e.OrganizationID)
}
func (f *Framework) CursorKey(orderBy FrameworkOrderField) page.CursorKey { func (f *Framework) CursorKey(orderBy FrameworkOrderField) page.CursorKey {
switch orderBy { switch orderBy {
case FrameworkOrderFieldCreatedAt: case FrameworkOrderFieldCreatedAt:
@@ -192,7 +167,7 @@ LIMIT 1;
framework, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[Framework]) framework, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[Framework])
if err != nil { if err != nil {
if errors.Is(err, pgx.ErrNoRows) { if errors.Is(err, pgx.ErrNoRows) {
return &ErrFrameworkNotFound{Identifier: referenceID} return ErrResourceNotFound
} }
return fmt.Errorf("cannot collect framework: %w", err) return fmt.Errorf("cannot collect framework: %w", err)
@@ -240,7 +215,7 @@ LIMIT 1;
framework, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[Framework]) framework, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[Framework])
if err != nil { if err != nil {
if errors.Is(err, pgx.ErrNoRows) { if errors.Is(err, pgx.ErrNoRows) {
return &ErrFrameworkNotFound{Identifier: frameworkID.String()} return ErrResourceNotFound
} }
return fmt.Errorf("cannot collect framework: %w", err) return fmt.Errorf("cannot collect framework: %w", err)
@@ -302,10 +277,7 @@ VALUES (
var pgErr *pgconn.PgError var pgErr *pgconn.PgError
if errors.As(err, &pgErr) { if errors.As(err, &pgErr) {
if pgErr.Code == "23505" && pgErr.ConstraintName == "frameworks_org_ref_unique" { if pgErr.Code == "23505" && pgErr.ConstraintName == "frameworks_org_ref_unique" {
return &ErrFrameworkReferenceIDAlreadyExists{ return ErrResourceAlreadyExists
ReferenceID: f.ReferenceID,
OrganizationID: f.OrganizationID,
}
} }
} }

View File

@@ -42,24 +42,8 @@ type (
} }
Invitations []*Invitation Invitations []*Invitation
InvitationData struct {
InvitationID gid.GID `json:"invitation_id"`
OrganizationID gid.GID `json:"organization_id"`
Email mail.Addr `json:"email"`
FullName string `json:"full_name"`
Role MembershipRole `json:"role"`
}
ErrInvitationNotFound struct {
ID string
}
) )
func (e ErrInvitationNotFound) Error() string {
return fmt.Sprintf("invitation not found: %s", e.ID)
}
func (i Invitation) CursorKey(orderBy InvitationOrderField) page.CursorKey { func (i Invitation) CursorKey(orderBy InvitationOrderField) page.CursorKey {
switch orderBy { switch orderBy {
case InvitationOrderFieldFullName: case InvitationOrderFieldFullName:
@@ -83,7 +67,7 @@ func (i Invitation) CursorKey(orderBy InvitationOrderField) page.CursorKey {
panic(fmt.Sprintf("unsupported order by: %s", orderBy)) panic(fmt.Sprintf("unsupported order by: %s", orderBy))
} }
func (i *Invitation) Create(ctx context.Context, conn pg.Conn, scope Scoper) error { func (i *Invitation) Insert(ctx context.Context, conn pg.Conn, scope Scoper) error {
query := ` query := `
INSERT INTO INSERT INTO
authz_invitations ( authz_invitations (
@@ -170,8 +154,9 @@ WHERE
invitation, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[Invitation]) invitation, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[Invitation])
if err != nil { if err != nil {
if errors.Is(err, pgx.ErrNoRows) { if errors.Is(err, pgx.ErrNoRows) {
return ErrInvitationNotFound{ID: id.String()} return ErrResourceNotFound
} }
return fmt.Errorf("cannot collect invitation: %w", err) return fmt.Errorf("cannot collect invitation: %w", err)
} }
@@ -204,25 +189,25 @@ WHERE
} }
if result.RowsAffected() == 0 { if result.RowsAffected() == 0 {
return ErrInvitationNotFound{ID: i.ID.String()} return ErrResourceNotFound
} }
return nil return nil
} }
func (i *Invitation) Delete(ctx context.Context, conn pg.Conn, scope Scoper) error { func (i *Invitation) Delete(ctx context.Context, conn pg.Conn, scope Scoper, invitationID gid.GID) error {
query := ` query := `
DELETE FROM DELETE FROM
authz_invitations authz_invitations
WHERE WHERE
id = @id %s
AND %s AND id = @invitation_id
` `
query = fmt.Sprintf(query, scope.SQLFragment()) query = fmt.Sprintf(query, scope.SQLFragment())
args := pgx.StrictNamedArgs{ args := pgx.StrictNamedArgs{
"id": i.ID, "invitation_id": invitationID,
} }
maps.Copy(args, scope.SQLArguments()) maps.Copy(args, scope.SQLArguments())
@@ -232,13 +217,13 @@ WHERE
} }
if result.RowsAffected() == 0 { if result.RowsAffected() == 0 {
return ErrInvitationNotFound{ID: i.ID.String()} return ErrResourceNotFound
} }
return nil return nil
} }
func (i *Invitations) LoadByEmail( func (i *Invitations) LoadByIdentityID(
ctx context.Context, ctx context.Context,
conn pg.Conn, conn pg.Conn,
scope Scoper, scope Scoper,
@@ -386,7 +371,7 @@ WHERE
func (i *Invitations) CountByEmail( func (i *Invitations) CountByEmail(
ctx context.Context, ctx context.Context,
conn pg.Conn, conn pg.Conn,
email string, email mail.Addr,
filter *InvitationFilter, filter *InvitationFilter,
) (int, error) { ) (int, error) {
q := ` q := `

View File

@@ -43,24 +43,8 @@ type (
} }
Measures []*Measure Measures []*Measure
ErrMeasureNotFound struct {
Identifier string
}
ErrMeasureAlreadyExists struct {
message string
}
) )
func (e ErrMeasureNotFound) Error() string {
return fmt.Sprintf("measure not found: %q", e.Identifier)
}
func (e ErrMeasureAlreadyExists) Error() string {
return e.message
}
func (m Measure) CursorKey(orderBy MeasureOrderField) page.CursorKey { func (m Measure) CursorKey(orderBy MeasureOrderField) page.CursorKey {
switch orderBy { switch orderBy {
case MeasureOrderFieldCreatedAt: case MeasureOrderFieldCreatedAt:
@@ -414,7 +398,7 @@ LIMIT 1;
measure, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[Measure]) measure, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[Measure])
if err != nil { if err != nil {
if errors.Is(err, pgx.ErrNoRows) { if errors.Is(err, pgx.ErrNoRows) {
return &ErrMeasureNotFound{Identifier: measureID.String()} return ErrResourceNotFound
} }
return fmt.Errorf("cannot collect measures: %w", err) return fmt.Errorf("cannot collect measures: %w", err)
@@ -552,9 +536,7 @@ VALUES (
var pgErr *pgconn.PgError var pgErr *pgconn.PgError
if errors.As(err, &pgErr) { if errors.As(err, &pgErr) {
if pgErr.Code == "23505" && pgErr.ConstraintName == "mitigations_org_ref_unique" { if pgErr.Code == "23505" && pgErr.ConstraintName == "mitigations_org_ref_unique" {
return &ErrMeasureAlreadyExists{ return ErrResourceAlreadyExists
message: fmt.Sprintf("measure with organization_id %s and reference_id %q already exists", m.OrganizationID, m.ReferenceID),
}
} }
} }
return fmt.Errorf("cannot insert measure: %w", err) return fmt.Errorf("cannot insert measure: %w", err)

View File

@@ -39,24 +39,8 @@ type (
} }
Meetings []*Meeting Meetings []*Meeting
ErrMeetingNotFound struct {
Identifier string
}
ErrMeetingAlreadyExists struct {
message string
}
) )
func (e ErrMeetingNotFound) Error() string {
return fmt.Sprintf("meeting not found: %s", e.Identifier)
}
func (e ErrMeetingAlreadyExists) Error() string {
return e.message
}
func (m Meeting) CursorKey(orderBy MeetingOrderField) page.CursorKey { func (m Meeting) CursorKey(orderBy MeetingOrderField) page.CursorKey {
switch orderBy { switch orderBy {
case MeetingOrderFieldCreatedAt: case MeetingOrderFieldCreatedAt:
@@ -106,7 +90,7 @@ LIMIT 1;
meeting, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[Meeting]) meeting, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[Meeting])
if err != nil { if err != nil {
if errors.Is(err, pgx.ErrNoRows) { if errors.Is(err, pgx.ErrNoRows) {
return &ErrMeetingNotFound{Identifier: meetingID.String()} return ErrResourceNotFound
} }
return fmt.Errorf("cannot collect meeting: %w", err) return fmt.Errorf("cannot collect meeting: %w", err)
@@ -271,7 +255,7 @@ WHERE %s
} }
if result.RowsAffected() == 0 { if result.RowsAffected() == 0 {
return &ErrMeetingNotFound{Identifier: m.ID.String()} return ErrResourceNotFound
} }
return nil return nil
@@ -300,7 +284,7 @@ WHERE %s
} }
if result.RowsAffected() == 0 { if result.RowsAffected() == 0 {
return &ErrMeetingNotFound{Identifier: m.ID.String()} return ErrResourceNotFound
} }
return nil return nil

View File

@@ -43,26 +43,8 @@ type (
} }
Memberships []*Membership Memberships []*Membership
ErrMembershipNotFound struct {
UserID gid.GID
OrgID gid.GID
}
ErrMembershipAlreadyExists struct {
UserID gid.GID
OrgID gid.GID
}
) )
func (e ErrMembershipNotFound) Error() string {
return fmt.Sprintf("membership not found for user %s in organization %s", e.UserID, e.OrgID)
}
func (e ErrMembershipAlreadyExists) Error() string {
return fmt.Sprintf("membership already exists for user %s in organization %s", e.UserID, e.OrgID)
}
func (m Membership) CursorKey(orderBy MembershipOrderField) page.CursorKey { func (m Membership) CursorKey(orderBy MembershipOrderField) page.CursorKey {
switch orderBy { switch orderBy {
case MembershipOrderFieldFullName: case MembershipOrderFieldFullName:
@@ -78,7 +60,46 @@ func (m Membership) CursorKey(orderBy MembershipOrderField) page.CursorKey {
panic(fmt.Sprintf("unsupported order by: %s", orderBy)) panic(fmt.Sprintf("unsupported order by: %s", orderBy))
} }
func (m *Membership) Create(ctx context.Context, conn pg.Conn, scope Scoper) error { func (m *Membership) LoadByUserInOrganization(ctx context.Context, conn pg.Conn, userID gid.GID, organizationID gid.GID) error {
q := `
SELECT
id,
user_id,
organization_id,
role,
created_at,
updated_at
FROM
authz_memberships
WHERE
user_id = @user_id
AND organization_id = @organization_id
`
args := pgx.StrictNamedArgs{
"user_id": userID,
"organization_id": organizationID,
}
rows, err := conn.Query(ctx, q, args)
if err != nil {
return fmt.Errorf("cannot query membership: %w", err)
}
membership, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[Membership])
if err != nil {
if errors.Is(err, pgx.ErrNoRows) {
return ErrResourceNotFound
}
return fmt.Errorf("cannot collect membership: %w", err)
}
*m = membership
return nil
}
func (m *Membership) Insert(ctx context.Context, conn pg.Conn, scope Scoper) error {
query := ` query := `
INSERT INTO INSERT INTO
authz_memberships ( authz_memberships (
@@ -115,8 +136,9 @@ VALUES (
if err != nil { if err != nil {
var pgErr *pgconn.PgError var pgErr *pgconn.PgError
if errors.As(err, &pgErr) && pgErr.Code == "23505" { if errors.As(err, &pgErr) && pgErr.Code == "23505" {
return ErrMembershipAlreadyExists{UserID: m.UserID, OrgID: m.OrganizationID} return ErrResourceAlreadyExists
} }
return fmt.Errorf("cannot create membership: %w", err) return fmt.Errorf("cannot create membership: %w", err)
} }
@@ -178,8 +200,9 @@ JOIN
membership, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[Membership]) membership, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[Membership])
if err != nil { if err != nil {
if errors.Is(err, pgx.ErrNoRows) { if errors.Is(err, pgx.ErrNoRows) {
return ErrMembershipNotFound{UserID: gid.GID{}, OrgID: gid.GID{}} return ErrResourceNotFound
} }
return fmt.Errorf("cannot collect membership: %w", err) return fmt.Errorf("cannot collect membership: %w", err)
} }
@@ -243,7 +266,7 @@ LIMIT 1;
defer rows.Close() defer rows.Close()
if !rows.Next() { if !rows.Next() {
return &ErrMembershipNotFound{UserID: userID, OrgID: entityID} return ErrResourceNotFound
} }
var membership Membership var membership Membership
@@ -269,9 +292,9 @@ func (m *Membership) LoadByUserAndOrg(
conn pg.Conn, conn pg.Conn,
scope Scoper, scope Scoper,
userID gid.GID, userID gid.GID,
orgID gid.GID, organizationID gid.GID,
) error { ) error {
query := ` q := `
WITH mbr AS ( WITH mbr AS (
SELECT SELECT
am.id, am.id,
@@ -302,20 +325,15 @@ JOIN
users u ON mbr.user_id = u.id users u ON mbr.user_id = u.id
` `
// Build scope fragment with table alias q = fmt.Sprintf(q, scope.SQLFragment())
scopeFragment := scope.SQLFragment()
// Replace column references with table-qualified versions
scopeFragment = strings.ReplaceAll(scopeFragment, "tenant_id =", "am.tenant_id =")
query = fmt.Sprintf(query, scopeFragment)
args := pgx.StrictNamedArgs{ args := pgx.StrictNamedArgs{
"user_id": userID, "user_id": userID,
"organization_id": orgID, "organization_id": organizationID,
} }
maps.Copy(args, scope.SQLArguments()) maps.Copy(args, scope.SQLArguments())
rows, err := conn.Query(ctx, query, args) rows, err := conn.Query(ctx, q, args)
if err != nil { if err != nil {
return fmt.Errorf("cannot query membership: %w", err) return fmt.Errorf("cannot query membership: %w", err)
} }
@@ -323,8 +341,9 @@ JOIN
membership, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[Membership]) membership, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[Membership])
if err != nil { if err != nil {
if errors.Is(err, pgx.ErrNoRows) { if errors.Is(err, pgx.ErrNoRows) {
return ErrMembershipNotFound{UserID: userID, OrgID: orgID} return ErrResourceNotFound
} }
return fmt.Errorf("cannot collect membership: %w", err) return fmt.Errorf("cannot collect membership: %w", err)
} }
@@ -359,25 +378,25 @@ WHERE
} }
if result.RowsAffected() == 0 { if result.RowsAffected() == 0 {
return ErrMembershipNotFound{UserID: m.UserID, OrgID: m.OrganizationID} return ErrResourceNotFound
} }
return nil return nil
} }
func (m *Membership) Delete(ctx context.Context, conn pg.Conn, scope Scoper) error { func (m *Membership) Delete(ctx context.Context, conn pg.Conn, scope Scoper, membershipID gid.GID) error {
query := ` query := `
DELETE FROM DELETE FROM
authz_memberships authz_memberships
WHERE WHERE
id = @id %s
AND %s AND id = @membership_id
` `
query = fmt.Sprintf(query, scope.SQLFragment()) query = fmt.Sprintf(query, scope.SQLFragment())
args := pgx.StrictNamedArgs{ args := pgx.StrictNamedArgs{
"id": m.ID, "membership_id": membershipID,
} }
maps.Copy(args, scope.SQLArguments()) maps.Copy(args, scope.SQLArguments())
@@ -387,7 +406,7 @@ WHERE
} }
if result.RowsAffected() == 0 { if result.RowsAffected() == 0 {
return ErrMembershipNotFound{UserID: m.UserID, OrgID: m.OrganizationID} return ErrResourceNotFound
} }
return nil return nil
@@ -398,6 +417,7 @@ func (m *Memberships) LoadByUserID(
conn pg.Conn, conn pg.Conn,
scope Scoper, scope Scoper,
userID gid.GID, userID gid.GID,
cursor *page.Cursor[MembershipOrderField],
) error { ) error {
query := ` query := `
WITH mbr AS ( WITH mbr AS (
@@ -552,3 +572,29 @@ WHERE
} }
return count, nil return count, nil
} }
func (m *Memberships) CountByUserID(
ctx context.Context,
conn pg.Conn,
userID gid.GID,
) (int, error) {
query := `
SELECT
COUNT(*)
FROM
authz_memberships
WHERE
user_id = @user_id
`
args := pgx.StrictNamedArgs{
"user_id": userID,
}
row := conn.QueryRow(ctx, query, args)
var count int
if err := row.Scan(&count); err != nil {
return 0, fmt.Errorf("cannot count memberships: %w", err)
}
return count, nil
}

View File

@@ -0,0 +1,9 @@
CREATE TYPE session_expire_reason AS ENUM (
'idle_timeout',
'revoked',
'closed'
);
ALTER TABLE sessions ADD COLUMN expire_reason session_expire_reason;
UPDATE sessions SET expire_reason = 'idle_timeout' WHERE expired_at < NOW();

View File

@@ -0,0 +1,2 @@
ALTER TABLE sessions ADD COLUMN user_agent TEXT DEFAULT 'SESSION_CREATED_BEFORE_USER_AGENT_COLUMN_ADDED';
ALTER TABLE sessions ADD COLUMN ip_address INET DEFAULT '::1';

View File

@@ -0,0 +1,8 @@
ALTER TABLE sessions ADD COLUMN tenant_id TEXT;
ALTER TABLE sessions ADD COLUMN parent_session_id TEXT REFERENCES sessions(id);
ALTER TABLE sessions ADD CONSTRAINT session_tenant_check CHECK (
(parent_session_id IS NULL AND tenant_id IS NULL) OR
(parent_session_id IS NOT NULL AND tenant_id IS NOT NULL)
);

View File

@@ -0,0 +1,2 @@
ALTER TABLE auth_user_api_keys ADD COLUMN expire_reason TEXT;

View File

@@ -44,24 +44,8 @@ type (
} }
Organizations []*Organization Organizations []*Organization
ErrOrganizationNotFound struct {
Identifier string
}
ErrOrganizationAlreadyExists struct {
message string
}
) )
func (e ErrOrganizationNotFound) Error() string {
return fmt.Sprintf("organization not found: %q", e.Identifier)
}
func (e ErrOrganizationAlreadyExists) Error() string {
return e.message
}
func (o Organization) CursorKey(orderBy OrganizationOrderField) page.CursorKey { func (o Organization) CursorKey(orderBy OrganizationOrderField) page.CursorKey {
switch orderBy { switch orderBy {
case OrganizationOrderFieldName: case OrganizationOrderFieldName:
@@ -116,7 +100,7 @@ LIMIT 1;
organization, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[Organization]) organization, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[Organization])
if err != nil { if err != nil {
if errors.Is(err, pgx.ErrNoRows) { if errors.Is(err, pgx.ErrNoRows) {
return &ErrOrganizationNotFound{Identifier: organizationID.String()} return ErrResourceNotFound
} }
return fmt.Errorf("cannot collect organization: %w", err) return fmt.Errorf("cannot collect organization: %w", err)
@@ -442,19 +426,14 @@ WHERE
func (o *Organization) Delete( func (o *Organization) Delete(
ctx context.Context, ctx context.Context,
conn pg.Conn, conn pg.Conn,
scope Scoper, organizationID gid.GID,
) error { ) error {
q := ` q := `
DELETE FROM organizations DELETE FROM organizations
WHERE WHERE id = @id
%s
AND id = @id
` `
q = fmt.Sprintf(q, scope.SQLFragment())
args := pgx.StrictNamedArgs{"id": o.ID} args := pgx.StrictNamedArgs{"id": o.ID}
maps.Copy(args, scope.SQLArguments())
_, err := conn.Exec(ctx, q, args) _, err := conn.Exec(ctx, q, args)
if err != nil { if err != nil {
@@ -505,7 +484,7 @@ LIMIT 1
organization, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[Organization]) organization, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[Organization])
if err != nil { if err != nil {
if errors.Is(err, pgx.ErrNoRows) { if errors.Is(err, pgx.ErrNoRows) {
return &ErrOrganizationNotFound{Identifier: customDomainID.String()} return ErrResourceNotFound
} }
return fmt.Errorf("cannot collect organization: %w", err) return fmt.Errorf("cannot collect organization: %w", err)

View File

@@ -33,16 +33,8 @@ type (
CreatedAt time.Time `db:"created_at"` CreatedAt time.Time `db:"created_at"`
UpdatedAt time.Time `db:"updated_at"` UpdatedAt time.Time `db:"updated_at"`
} }
ErrOrganizationContextNotFound struct {
Identifier string
}
) )
func (e ErrOrganizationContextNotFound) Error() string {
return fmt.Sprintf("organization context not found: %q", e.Identifier)
}
func (oc *OrganizationContext) LoadByOrganizationID( func (oc *OrganizationContext) LoadByOrganizationID(
ctx context.Context, ctx context.Context,
conn pg.Conn, conn pg.Conn,
@@ -76,7 +68,7 @@ LIMIT 1;
orgContext, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[OrganizationContext]) orgContext, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[OrganizationContext])
if err != nil { if err != nil {
if errors.Is(err, pgx.ErrNoRows) { if errors.Is(err, pgx.ErrNoRows) {
return &ErrOrganizationContextNotFound{Identifier: organizationID.String()} return ErrResourceNotFound
} }
return fmt.Errorf("cannot collect organization context: %w", err) return fmt.Errorf("cannot collect organization context: %w", err)
@@ -155,7 +147,7 @@ WHERE
} }
if result.RowsAffected() == 0 { if result.RowsAffected() == 0 {
return &ErrOrganizationContextNotFound{Identifier: oc.OrganizationID.String()} return ErrResourceNotFound
} }
return nil return nil

View File

@@ -45,32 +45,8 @@ type (
} }
Peoples []*People Peoples []*People
ErrPeopleNotFound struct {
Identifier string
}
ErrPeopleAlreadyExists struct {
message string
}
ErrPeopleReferenced struct {
message string
}
) )
func (e ErrPeopleNotFound) Error() string {
return fmt.Sprintf("people not found: %s", e.Identifier)
}
func (e ErrPeopleAlreadyExists) Error() string {
return e.message
}
func (e ErrPeopleReferenced) Error() string {
return e.message
}
func (p People) CursorKey(orderBy PeopleOrderField) page.CursorKey { func (p People) CursorKey(orderBy PeopleOrderField) page.CursorKey {
switch orderBy { switch orderBy {
case PeopleOrderFieldCreatedAt: case PeopleOrderFieldCreatedAt:
@@ -124,7 +100,7 @@ LIMIT 1;
people, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[People]) people, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[People])
if err != nil { if err != nil {
if errors.Is(err, pgx.ErrNoRows) { if errors.Is(err, pgx.ErrNoRows) {
return &ErrPeopleNotFound{Identifier: peopleID.String()} return ErrResourceNotFound
} }
return fmt.Errorf("cannot collect people: %w", err) return fmt.Errorf("cannot collect people: %w", err)
@@ -175,7 +151,7 @@ LIMIT 1;
people, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[People]) people, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[People])
if err != nil { if err != nil {
if errors.Is(err, pgx.ErrNoRows) { if errors.Is(err, pgx.ErrNoRows) {
return &ErrPeopleNotFound{Identifier: primaryEmailAddress} return ErrResourceNotFound
} }
return fmt.Errorf("cannot collect people: %w", err) return fmt.Errorf("cannot collect people: %w", err)
@@ -231,7 +207,7 @@ LIMIT 1;
people, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[People]) people, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[People])
if err != nil { if err != nil {
if errors.Is(err, pgx.ErrNoRows) { if errors.Is(err, pgx.ErrNoRows) {
return &ErrPeopleNotFound{Identifier: primaryEmailAddress.String()} return ErrResourceNotFound
} }
return fmt.Errorf("cannot collect people: %w", err) return fmt.Errorf("cannot collect people: %w", err)
@@ -362,9 +338,7 @@ DELETE FROM peoples WHERE %s AND id = @people_id
var pgErr *pgconn.PgError var pgErr *pgconn.PgError
if errors.As(err, &pgErr) { if errors.As(err, &pgErr) {
if pgErr.Code == "23503" { if pgErr.Code == "23503" {
return &ErrPeopleReferenced{ return ErrResourceInUse
message: fmt.Sprintf("person with id %s cannot be deleted because it is referenced by other records", p.ID),
}
} }
} }
return fmt.Errorf("cannot delete person: %w", err) return fmt.Errorf("cannot delete person: %w", err)

View File

@@ -79,7 +79,7 @@ LIMIT 1;
report, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[Report]) report, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[Report])
if err != nil { if err != nil {
if errors.Is(err, pgx.ErrNoRows) { if errors.Is(err, pgx.ErrNoRows) {
return &ErrDocumentNotFound{Identifier: reportID.String()} return ErrResourceNotFound
} }
return fmt.Errorf("cannot collect report: %w", err) return fmt.Errorf("cannot collect report: %w", err)

View File

@@ -57,24 +57,8 @@ type (
RiskSnapshotter interface { RiskSnapshotter interface {
InsertRiskSnapshots(ctx context.Context, conn pg.Conn, scope Scoper, organizationID, snapshotID gid.GID) error InsertRiskSnapshots(ctx context.Context, conn pg.Conn, scope Scoper, organizationID, snapshotID gid.GID) error
} }
ErrRiskNotFound struct {
Identifier string
}
ErrRiskAlreadyExists struct {
message string
}
) )
func (e ErrRiskNotFound) Error() string {
return fmt.Sprintf("risk not found: %q", e.Identifier)
}
func (e ErrRiskAlreadyExists) Error() string {
return e.message
}
func (r *Risk) CursorKey(orderBy RiskOrderField) page.CursorKey { func (r *Risk) CursorKey(orderBy RiskOrderField) page.CursorKey {
switch orderBy { switch orderBy {
case RiskOrderFieldCreatedAt: case RiskOrderFieldCreatedAt:
@@ -392,7 +376,7 @@ LIMIT 1;
risk, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[Risk]) risk, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[Risk])
if err != nil { if err != nil {
if errors.Is(err, pgx.ErrNoRows) { if errors.Is(err, pgx.ErrNoRows) {
return &ErrRiskNotFound{Identifier: riskID.String()} return ErrResourceNotFound
} }
return fmt.Errorf("cannot collect risk: %w", err) return fmt.Errorf("cannot collect risk: %w", err)

View File

@@ -16,12 +16,14 @@ package coredata
import ( import (
"context" "context"
"errors"
"fmt" "fmt"
"time" "time"
"go.probo.inc/probo/pkg/gid"
"github.com/jackc/pgx/v5" "github.com/jackc/pgx/v5"
"github.com/jackc/pgx/v5/pgconn"
"go.gearno.de/kit/pg" "go.gearno.de/kit/pg"
"go.probo.inc/probo/pkg/gid"
) )
type SAMLAssertion struct { type SAMLAssertion struct {
@@ -31,27 +33,17 @@ type SAMLAssertion struct {
ExpiresAt time.Time `db:"expires_at"` ExpiresAt time.Time `db:"expires_at"`
} }
type ErrAssertionAlreadyUsed struct {
AssertionID string
}
func (e ErrAssertionAlreadyUsed) Error() string {
return fmt.Sprintf("assertion ID %q has already been used (replay attack)", e.AssertionID)
}
func (s *SAMLAssertion) Insert( func (s *SAMLAssertion) Insert(
ctx context.Context, ctx context.Context,
conn pg.Conn, conn pg.Conn,
scope Scoper,
) error { ) error {
query := ` query := `
INSERT INTO auth_saml_assertions (id, tenant_id, organization_id, used_at, expires_at) INSERT INTO auth_saml_assertions (id, organization_id, used_at, expires_at)
VALUES (@id, @tenant_id, @organization_id, @used_at, @expires_at) VALUES (@id, @organization_id, @used_at, @expires_at)
` `
args := pgx.NamedArgs{ args := pgx.NamedArgs{
"id": s.ID, "id": s.ID,
"tenant_id": scope.GetTenantID(),
"organization_id": s.OrganizationID, "organization_id": s.OrganizationID,
"used_at": s.UsedAt, "used_at": s.UsedAt,
"expires_at": s.ExpiresAt, "expires_at": s.ExpiresAt,
@@ -59,6 +51,11 @@ VALUES (@id, @tenant_id, @organization_id, @used_at, @expires_at)
_, err := conn.Exec(ctx, query, args) _, err := conn.Exec(ctx, query, args)
if err != nil { if err != nil {
var pgErr *pgconn.PgError
if errors.As(err, &pgErr) && pgErr.Code == "23505" && pgErr.ConstraintName == "auth_saml_assertions_pkey" {
return ErrResourceAlreadyExists
}
return fmt.Errorf("cannot insert saml_assertion: %w", err) return fmt.Errorf("cannot insert saml_assertion: %w", err)
} }

View File

@@ -16,6 +16,8 @@ package coredata
import ( import (
"context" "context"
"crypto/x509"
"encoding/pem"
"fmt" "fmt"
"maps" "maps"
"time" "time"
@@ -23,28 +25,55 @@ import (
"github.com/jackc/pgx/v5" "github.com/jackc/pgx/v5"
"go.gearno.de/kit/pg" "go.gearno.de/kit/pg"
"go.probo.inc/probo/pkg/gid" "go.probo.inc/probo/pkg/gid"
"go.probo.inc/probo/pkg/page"
) )
type SAMLConfiguration struct { type (
ID gid.GID `db:"id"` SAMLConfiguration struct {
OrganizationID gid.GID `db:"organization_id"` ID gid.GID `db:"id"`
EmailDomain string `db:"email_domain"` OrganizationID gid.GID `db:"organization_id"`
Enabled bool `db:"enabled"` EmailDomain string `db:"email_domain"`
EnforcementPolicy SAMLEnforcementPolicy `db:"enforcement_policy"` EnforcementPolicy SAMLEnforcementPolicy `db:"enforcement_policy"`
IdPEntityID string `db:"idp_entity_id"` IdPEntityID string `db:"idp_entity_id"`
IdPSsoURL string `db:"idp_sso_url"` IdPSsoURL string `db:"idp_sso_url"`
IdPCertificate string `db:"idp_certificate"` IdPCertificate string `db:"idp_certificate"`
IdPMetadataURL *string `db:"idp_metadata_url"` IdPMetadataURL *string `db:"idp_metadata_url"`
AttributeEmail string `db:"attribute_email"` AttributeEmail string `db:"attribute_email"`
AttributeFirstname string `db:"attribute_firstname"` AttributeFirstname string `db:"attribute_firstname"`
AttributeLastname string `db:"attribute_lastname"` AttributeLastname string `db:"attribute_lastname"`
AttributeRole string `db:"attribute_role"` AttributeRole string `db:"attribute_role"`
AutoSignupEnabled bool `db:"auto_signup_enabled"` AutoSignupEnabled bool `db:"auto_signup_enabled"`
DomainVerified bool `db:"domain_verified"` DomainVerified bool `db:"domain_verified"`
DomainVerificationToken *string `db:"domain_verification_token"` DomainVerificationToken *string `db:"domain_verification_token"`
DomainVerifiedAt *time.Time `db:"domain_verified_at"` DomainVerifiedAt *time.Time `db:"domain_verified_at"`
CreatedAt time.Time `db:"created_at"` CreatedAt time.Time `db:"created_at"`
UpdatedAt time.Time `db:"updated_at"` UpdatedAt time.Time `db:"updated_at"`
}
SAMLConfigurations []*SAMLConfiguration
)
func (s *SAMLConfiguration) CursorKey(orderBy SAMLConfigurationOrderField) page.CursorKey {
switch orderBy {
case SAMLConfigurationOrderFieldCreatedAt:
return page.NewCursorKey(s.ID, s.CreatedAt)
}
panic(fmt.Sprintf("unsupported order by: %s", orderBy))
}
func (s *SAMLConfiguration) GetIdPCertificate() (*x509.Certificate, error) {
block, _ := pem.Decode([]byte(s.IdPCertificate))
if block == nil {
return nil, fmt.Errorf("cannot decode PEM block from IdP certificate")
}
cert, err := x509.ParseCertificate(block.Bytes)
if err != nil {
return nil, fmt.Errorf("cannot parse X.509 certificate: %w", err)
}
return cert, nil
} }
func (s *SAMLConfiguration) LoadByOrganizationIDAndEmailDomain( func (s *SAMLConfiguration) LoadByOrganizationIDAndEmailDomain(
@@ -154,6 +183,10 @@ LIMIT 1;
config, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[SAMLConfiguration]) config, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[SAMLConfiguration])
if err != nil { if err != nil {
if err == pgx.ErrNoRows {
return ErrResourceNotFound
}
return fmt.Errorf("cannot collect saml_configuration: %w", err) return fmt.Errorf("cannot collect saml_configuration: %w", err)
} }
@@ -173,7 +206,6 @@ INSERT INTO auth_saml_configurations (
tenant_id, tenant_id,
organization_id, organization_id,
email_domain, email_domain,
enabled,
enforcement_policy, enforcement_policy,
idp_entity_id, idp_entity_id,
idp_sso_url, idp_sso_url,
@@ -194,7 +226,6 @@ INSERT INTO auth_saml_configurations (
@tenant_id, @tenant_id,
@organization_id, @organization_id,
@email_domain, @email_domain,
@enabled,
@enforcement_policy, @enforcement_policy,
@idp_entity_id, @idp_entity_id,
@idp_sso_url, @idp_sso_url,
@@ -218,7 +249,6 @@ INSERT INTO auth_saml_configurations (
"tenant_id": scope.GetTenantID(), "tenant_id": scope.GetTenantID(),
"organization_id": s.OrganizationID, "organization_id": s.OrganizationID,
"email_domain": s.EmailDomain, "email_domain": s.EmailDomain,
"enabled": s.Enabled,
"enforcement_policy": s.EnforcementPolicy, "enforcement_policy": s.EnforcementPolicy,
"idp_entity_id": s.IdPEntityID, "idp_entity_id": s.IdPEntityID,
"idp_sso_url": s.IdPSsoURL, "idp_sso_url": s.IdPSsoURL,
@@ -252,7 +282,6 @@ func (s *SAMLConfiguration) Update(
q := ` q := `
UPDATE auth_saml_configurations UPDATE auth_saml_configurations
SET SET
enabled = @enabled,
enforcement_policy = @enforcement_policy, enforcement_policy = @enforcement_policy,
idp_entity_id = @idp_entity_id, idp_entity_id = @idp_entity_id,
idp_sso_url = @idp_sso_url, idp_sso_url = @idp_sso_url,
@@ -276,7 +305,6 @@ WHERE
args := pgx.StrictNamedArgs{ args := pgx.StrictNamedArgs{
"id": s.ID, "id": s.ID,
"enabled": s.Enabled,
"enforcement_policy": s.EnforcementPolicy, "enforcement_policy": s.EnforcementPolicy,
"idp_entity_id": s.IdPEntityID, "idp_entity_id": s.IdPEntityID,
"idp_sso_url": s.IdPSsoURL, "idp_sso_url": s.IdPSsoURL,
@@ -328,18 +356,17 @@ WHERE
return nil return nil
} }
func LoadSAMLConfigurationsByOrganizationID( func (s *SAMLConfigurations) LoadByOrganizationID(
ctx context.Context, ctx context.Context,
conn pg.Conn, conn pg.Conn,
scope Scoper, scope Scoper,
organizationID gid.GID, organizationID gid.GID,
) ([]*SAMLConfiguration, error) { ) error {
q := ` q := `
SELECT SELECT
id, id,
organization_id, organization_id,
email_domain, email_domain,
enabled,
enforcement_policy, enforcement_policy,
idp_entity_id, idp_entity_id,
idp_sso_url, idp_sso_url,
@@ -370,20 +397,17 @@ ORDER BY email_domain ASC;
rows, err := conn.Query(ctx, q, args) rows, err := conn.Query(ctx, q, args)
if err != nil { if err != nil {
return nil, fmt.Errorf("cannot query auth_saml_configurations: %w", err) return fmt.Errorf("cannot query auth_saml_configurations: %w", err)
} }
configs, err := pgx.CollectRows(rows, pgx.RowToStructByName[SAMLConfiguration]) samlConfigurations, err := pgx.CollectRows(rows, pgx.RowToAddrOfStructByName[SAMLConfiguration])
if err != nil { if err != nil {
return nil, fmt.Errorf("cannot collect saml_configurations: %w", err) return fmt.Errorf("cannot collect saml_configurations: %w", err)
} }
result := make([]*SAMLConfiguration, len(configs)) *s = samlConfigurations
for i := range configs {
result[i] = &configs[i]
}
return result, nil return nil
} }
// LoadAllEnabledSAMLConfigurationsByEmailDomain loads all enabled SAML configurations for a given email domain // LoadAllEnabledSAMLConfigurationsByEmailDomain loads all enabled SAML configurations for a given email domain
@@ -398,7 +422,6 @@ SELECT
id, id,
organization_id, organization_id,
email_domain, email_domain,
enabled,
enforcement_policy, enforcement_policy,
idp_entity_id, idp_entity_id,
idp_sso_url, idp_sso_url,
@@ -458,7 +481,6 @@ SELECT
id, id,
organization_id, organization_id,
email_domain, email_domain,
enabled,
enforcement_policy, enforcement_policy,
idp_entity_id, idp_entity_id,
idp_sso_url, idp_sso_url,
@@ -503,3 +525,38 @@ WHERE
return result, nil return result, nil
} }
func (s *SAMLConfigurations) CountByOrganizationID(
ctx context.Context,
conn pg.Conn,
scope Scoper,
organizationID gid.GID,
) (int, error) {
q := `
SELECT
COUNT(*)
FROM
auth_saml_configurations
WHERE
%s
AND organization_id = @organization_id
`
q = fmt.Sprintf(q, scope.SQLFragment())
args := pgx.StrictNamedArgs{"organization_id": organizationID}
maps.Copy(args, scope.SQLArguments())
rows, err := conn.Query(ctx, q, args)
if err != nil {
return 0, fmt.Errorf("cannot query auth_saml_configurations: %w", err)
}
var count int
err = rows.Scan(&count)
if err != nil {
return 0, fmt.Errorf("cannot collect count: %w", err)
}
return count, nil
}

View File

@@ -0,0 +1,40 @@
// Copyright (c) 2025 Probo Inc <hello@getprobo.com>.
//
// Permission to use, copy, modify, and/or distribute this software for any
// purpose with or without fee is hereby granted, provided that the above
// copyright notice and this permission notice appear in all copies.
//
// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH
// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY
// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT,
// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM
// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR
// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR
// PERFORMANCE OF THIS SOFTWARE.
package coredata
type (
SAMLConfigurationOrderField string
)
const (
SAMLConfigurationOrderFieldCreatedAt SAMLConfigurationOrderField = "CREATED_AT"
)
func (p SAMLConfigurationOrderField) Column() string {
return string(p)
}
func (p SAMLConfigurationOrderField) String() string {
return string(p)
}
func (p SAMLConfigurationOrderField) MarshalText() ([]byte, error) {
return []byte(p.String()), nil
}
func (p *SAMLConfigurationOrderField) UnmarshalText(text []byte) error {
*p = SAMLConfigurationOrderField(text)
return nil
}

View File

@@ -19,9 +19,9 @@ import (
"fmt" "fmt"
"time" "time"
"go.probo.inc/probo/pkg/gid"
"github.com/jackc/pgx/v5" "github.com/jackc/pgx/v5"
"go.gearno.de/kit/pg" "go.gearno.de/kit/pg"
"go.probo.inc/probo/pkg/gid"
) )
type SAMLRequest struct { type SAMLRequest struct {
@@ -31,37 +31,18 @@ type SAMLRequest struct {
ExpiresAt time.Time `db:"expires_at"` ExpiresAt time.Time `db:"expires_at"`
} }
type ErrSAMLRequestNotFound struct {
RequestID string
}
func (e ErrSAMLRequestNotFound) Error() string {
return fmt.Sprintf("SAML request ID %q not found", e.RequestID)
}
type ErrSAMLRequestExpired struct {
RequestID string
ExpiresAt time.Time
}
func (e ErrSAMLRequestExpired) Error() string {
return fmt.Sprintf("SAML request ID %q expired at %v", e.RequestID, e.ExpiresAt)
}
func (s *SAMLRequest) Insert( func (s *SAMLRequest) Insert(
ctx context.Context, ctx context.Context,
conn pg.Conn, conn pg.Conn,
scope Scoper,
) error { ) error {
query := ` query := `
INSERT INTO auth_saml_requests (id, organization_id, tenant_id, created_at, expires_at) INSERT INTO auth_saml_requests (id, organization_id, created_at, expires_at)
VALUES (@id, @organization_id, @tenant_id, @created_at, @expires_at) VALUES (@id, @organization_id, @created_at, @expires_at)
` `
args := pgx.NamedArgs{ args := pgx.NamedArgs{
"id": s.ID, "id": s.ID,
"organization_id": s.OrganizationID, "organization_id": s.OrganizationID,
"tenant_id": scope.GetTenantID(),
"created_at": s.CreatedAt, "created_at": s.CreatedAt,
"expires_at": s.ExpiresAt, "expires_at": s.ExpiresAt,
} }
@@ -99,8 +80,9 @@ LIMIT 1
req, err := pgx.CollectOneRow(rows, pgx.RowToStructByName[SAMLRequest]) req, err := pgx.CollectOneRow(rows, pgx.RowToStructByName[SAMLRequest])
if err == pgx.ErrNoRows { if err == pgx.ErrNoRows {
return ErrSAMLRequestNotFound{RequestID: requestID} return ErrResourceNotFound
} }
if err != nil { if err != nil {
return fmt.Errorf("cannot collect saml_request: %w", err) return fmt.Errorf("cannot collect saml_request: %w", err)
} }

View File

@@ -17,8 +17,8 @@ package coredata
import ( import (
"fmt" "fmt"
"go.probo.inc/probo/pkg/gid"
"github.com/jackc/pgx/v5" "github.com/jackc/pgx/v5"
"go.probo.inc/probo/pkg/gid"
) )
type ( type (
@@ -62,6 +62,10 @@ func NewScope(tenantID gid.TenantID) *Scope {
} }
} }
func NewScopeFromObjectID(objectID gid.GID) *Scope {
return NewScope(objectID.TenantID())
}
func (s *Scope) SQLArguments() pgx.StrictNamedArgs { func (s *Scope) SQLArguments() pgx.StrictNamedArgs {
return pgx.StrictNamedArgs{ return pgx.StrictNamedArgs{
"tenant_id": s.tenantID, "tenant_id": s.tenantID,

View File

@@ -18,26 +18,35 @@ import (
"context" "context"
"errors" "errors"
"fmt" "fmt"
"maps"
"net"
"time" "time"
"go.probo.inc/probo/pkg/gid"
"go.probo.inc/probo/pkg/page"
"github.com/jackc/pgx/v5" "github.com/jackc/pgx/v5"
"go.gearno.de/kit/pg" "go.gearno.de/kit/pg"
"go.probo.inc/probo/pkg/gid"
"go.probo.inc/probo/pkg/page"
) )
type ( type (
Session struct { Session struct {
ID gid.GID `db:"id"` ID gid.GID `db:"id"`
UserID gid.GID `db:"user_id"` UserID gid.GID `db:"user_id"`
Data SessionData `db:"data"` TenantID *gid.TenantID `db:"tenant_id"`
ExpiredAt time.Time `db:"expired_at"` ParentSessionID *gid.GID `db:"parent_session_id"`
CreatedAt time.Time `db:"created_at"` Data SessionData `db:"data"`
UpdatedAt time.Time `db:"updated_at"` UserAgent string `db:"user_agent"`
IPAddress net.IP `db:"ip_address"`
ExpireReason *ExpireReason `db:"expire_reason"`
ExpiredAt time.Time `db:"expired_at"`
CreatedAt time.Time `db:"created_at"`
UpdatedAt time.Time `db:"updated_at"`
} }
Sessions []*Session
SessionData struct { SessionData struct {
PasswordAuthenticated bool `json:"password_authenticated"` PasswordAuthenticated bool `json:"password_authenticated"`
SAMLAuthenticatedOrgs map[string]SAMLAuthInfo `json:"saml_authenticated_orgs,omitempty"` SAMLAuthenticatedOrgs map[string]SAMLAuthInfo `json:"saml_authenticated_orgs,omitempty"`
} }
@@ -46,33 +55,39 @@ type (
SAMLConfigID gid.GID `json:"saml_config_id"` SAMLConfigID gid.GID `json:"saml_config_id"`
SAMLSubject string `json:"saml_subject"` SAMLSubject string `json:"saml_subject"`
} }
ErrSessionNotFound struct {
Identifier string
}
ErrSessionAlreadyExists struct {
message string
}
) )
func (e ErrSessionNotFound) Error() string { func NewRootSession(userID gid.GID, duration time.Duration) *Session {
return fmt.Sprintf("session not found: %q", e.Identifier) return &Session{
} ID: gid.New(gid.NilTenant, SessionEntityType),
UserID: userID,
func (e ErrSessionAlreadyExists) Error() string { ExpiredAt: time.Now().Add(duration),
return e.message CreatedAt: time.Now(),
UpdatedAt: time.Now(),
}
} }
func (s Session) CursorKey(orderBy SessionOrderField) page.CursorKey { func (s Session) CursorKey(orderBy SessionOrderField) page.CursorKey {
switch orderBy { switch orderBy {
case SessionOrderFieldCreatedAt: case SessionOrderFieldCreatedAt:
return page.NewCursorKey(s.ID, s.CreatedAt) return page.NewCursorKey(s.ID, s.CreatedAt)
case SessionOrderFieldExpiredAt:
return page.NewCursorKey(s.ID, s.ExpiredAt)
case SessionOrderFieldUpdatedAt:
return page.NewCursorKey(s.ID, s.UpdatedAt)
} }
panic(fmt.Sprintf("unsupported order by: %s", orderBy)) panic(fmt.Sprintf("unsupported order by: %s", orderBy))
} }
func (s *Session) IsRootSession() bool {
return s.ParentSessionID == nil
}
func (s *Session) IsChildSession() bool {
return s.ParentSessionID != nil
}
func (s *Session) LoadByID( func (s *Session) LoadByID(
ctx context.Context, ctx context.Context,
conn pg.Conn, conn pg.Conn,
@@ -82,7 +97,12 @@ func (s *Session) LoadByID(
SELECT SELECT
id, id,
user_id, user_id,
data, tenant_id,
data,
parent_session_id,
expire_reason,
user_agent,
ip_address,
expired_at, expired_at,
created_at, created_at,
updated_at updated_at
@@ -103,7 +123,7 @@ LIMIT 1;
session, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[Session]) session, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[Session])
if err != nil { if err != nil {
if errors.Is(err, pgx.ErrNoRows) { if errors.Is(err, pgx.ErrNoRows) {
return &ErrSessionNotFound{Identifier: sessionID.String()} return ErrResourceNotFound
} }
return fmt.Errorf("cannot collect session: %w", err) return fmt.Errorf("cannot collect session: %w", err)
@@ -119,11 +139,16 @@ func (s *Session) Insert(
) error { ) error {
q := ` q := `
INSERT INTO INSERT INTO
sessions (id, user_id, data, expired_at, created_at, updated_at) sessions (id, user_id, tenant_id, data, parent_session_id, expire_reason, user_agent, ip_address, expired_at, created_at, updated_at)
VALUES ( VALUES (
@session_id, @session_id,
@user_id, @user_id,
@tenant_id,
@data, @data,
@parent_session_id,
@expire_reason,
@user_agent,
@ip_address,
@expired_at, @expired_at,
@created_at, @created_at,
@updated_at @updated_at
@@ -131,12 +156,17 @@ VALUES (
` `
args := pgx.StrictNamedArgs{ args := pgx.StrictNamedArgs{
"session_id": s.ID, "session_id": s.ID,
"user_id": s.UserID, "user_id": s.UserID,
"data": s.Data, "tenant_id": s.TenantID,
"expired_at": s.ExpiredAt, "data": s.Data,
"created_at": s.CreatedAt, "parent_session_id": s.ParentSessionID,
"updated_at": s.UpdatedAt, "expire_reason": s.ExpireReason,
"user_agent": s.UserAgent,
"ip_address": s.IPAddress,
"expired_at": s.ExpiredAt,
"created_at": s.CreatedAt,
"updated_at": s.UpdatedAt,
} }
_, err := conn.Exec(ctx, q, args) _, err := conn.Exec(ctx, q, args)
@@ -152,36 +182,121 @@ UPDATE sessions
SET SET
expired_at = @expired_at, expired_at = @expired_at,
updated_at = @updated_at, updated_at = @updated_at,
user_agent = @user_agent,
ip_address = @ip_address,
expire_reason = @expire_reason,
data = @data data = @data
WHERE WHERE
id = @session_id id = @session_id
` `
args := pgx.StrictNamedArgs{ args := pgx.StrictNamedArgs{
"session_id": s.ID, "session_id": s.ID,
"data": s.Data, "user_agent": s.UserAgent,
"expired_at": s.ExpiredAt, "ip_address": s.IPAddress,
"updated_at": s.UpdatedAt, "expire_reason": s.ExpireReason,
"data": s.Data,
"expired_at": s.ExpiredAt,
"updated_at": s.UpdatedAt,
} }
_, err := conn.Exec(ctx, q, args) result, err := conn.Exec(ctx, q, args)
return err if err != nil {
return fmt.Errorf("cannot update session: %w", err)
}
if result.RowsAffected() == 0 {
return ErrResourceNotFound
}
return nil
} }
func DeleteSession( func (s *Sessions) LoadByUserID(ctx context.Context, conn pg.Conn, userID gid.GID, cursor *page.Cursor[SessionOrderField]) error {
ctx context.Context,
conn pg.Conn,
sessionID gid.GID,
) error {
q := ` q := `
DELETE FROM SELECT
id,
user_id,
tenant_id,
data,
parent_session_id,
expire_reason,
user_agent,
ip_address,
expired_at,
created_at,
updated_at
FROM
sessions sessions
WHERE WHERE
id = @session_id user_id = @user_id
AND %s
` `
args := pgx.StrictNamedArgs{"session_id": sessionID} q = fmt.Sprintf(q, cursor.SQLFragment())
_, err := conn.Exec(ctx, q, args) args := pgx.StrictNamedArgs{"user_id": userID}
return err maps.Copy(args, cursor.SQLArguments())
rows, err := conn.Query(ctx, q, args)
if err != nil {
return fmt.Errorf("cannot query sessions: %w", err)
}
sessions, err := pgx.CollectRows(rows, pgx.RowToAddrOfStructByName[Session])
if err != nil {
return fmt.Errorf("cannot collect sessions: %w", err)
}
*s = sessions
return nil
}
func (s *Sessions) CountByUserID(ctx context.Context, conn pg.Conn, userID gid.GID) (int, error) {
q := `
SELECT
COUNT(*)
FROM
sessions
WHERE
user_id = @user_id
`
args := pgx.StrictNamedArgs{"user_id": userID}
row := conn.QueryRow(ctx, q, args)
var count int
if err := row.Scan(&count); err != nil {
return 0, fmt.Errorf("cannot scan count: %w", err)
}
return count, nil
}
func (s *Sessions) ExpireAllForUserExceptOneSession(ctx context.Context, conn pg.Conn, userID gid.GID, sessionID gid.GID) (int64, error) {
q := `
UPDATE sessions
SET
expired_at = NOW(),
updated_at = NOW(),
expire_reason = 'revoked'
WHERE
id != @session_id
AND user_id = @user_id
AND expire_reason IS NULL
`
args := pgx.StrictNamedArgs{
"session_id": sessionID,
"user_id": userID,
}
result, err := conn.Exec(ctx, q, args)
if err != nil {
return 0, fmt.Errorf("cannot query sessions: %w", err)
}
return result.RowsAffected(), nil
} }

View File

@@ -20,9 +20,20 @@ type (
const ( const (
SessionOrderFieldCreatedAt SessionOrderField = "CREATED_AT" SessionOrderFieldCreatedAt SessionOrderField = "CREATED_AT"
SessionOrderFieldExpiredAt SessionOrderField = "EXPIRED_AT"
SessionOrderFieldUpdatedAt SessionOrderField = "UPDATED_AT"
) )
func (p SessionOrderField) Column() string { func (p SessionOrderField) Column() string {
switch p {
case SessionOrderFieldCreatedAt:
return "created_at"
case SessionOrderFieldExpiredAt:
return "expired_at"
case SessionOrderFieldUpdatedAt:
return "updated_at"
}
return string(p) return string(p)
} }

View File

@@ -41,24 +41,8 @@ type (
} }
StatesOfApplicability []*StateOfApplicability StatesOfApplicability []*StateOfApplicability
ErrStateOfApplicabilityNotFound struct {
Identifier string
}
ErrStateOfApplicabilityAlreadyExists struct {
message string
}
) )
func (e ErrStateOfApplicabilityNotFound) Error() string {
return fmt.Sprintf("state of applicability not found: %s", e.Identifier)
}
func (e ErrStateOfApplicabilityAlreadyExists) Error() string {
return e.message
}
func (s StateOfApplicability) CursorKey(orderBy StateOfApplicabilityOrderField) page.CursorKey { func (s StateOfApplicability) CursorKey(orderBy StateOfApplicabilityOrderField) page.CursorKey {
switch orderBy { switch orderBy {
case StateOfApplicabilityOrderFieldCreatedAt: case StateOfApplicabilityOrderFieldCreatedAt:
@@ -107,7 +91,7 @@ LIMIT 1;
stateOfApplicability, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[StateOfApplicability]) stateOfApplicability, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[StateOfApplicability])
if err != nil { if err != nil {
if errors.Is(err, pgx.ErrNoRows) { if errors.Is(err, pgx.ErrNoRows) {
return &ErrStateOfApplicabilityNotFound{Identifier: stateOfApplicabilityID.String()} return ErrResourceNotFound
} }
return fmt.Errorf("cannot collect state_of_applicability: %w", err) return fmt.Errorf("cannot collect state_of_applicability: %w", err)
@@ -246,9 +230,7 @@ VALUES (
var pgErr *pgconn.PgError var pgErr *pgconn.PgError
if errors.As(err, &pgErr) { if errors.As(err, &pgErr) {
if pgErr.Code == "23505" { if pgErr.Code == "23505" {
return &ErrStateOfApplicabilityAlreadyExists{ return ErrResourceAlreadyExists
message: fmt.Sprintf("state of applicability with name %q already exists", s.Name),
}
} }
} }
return fmt.Errorf("cannot insert state_of_applicability: %w", err) return fmt.Errorf("cannot insert state_of_applicability: %w", err)
@@ -287,16 +269,14 @@ WHERE %s
var pgErr *pgconn.PgError var pgErr *pgconn.PgError
if errors.As(err, &pgErr) { if errors.As(err, &pgErr) {
if pgErr.Code == "23505" { if pgErr.Code == "23505" {
return &ErrStateOfApplicabilityAlreadyExists{ return ErrResourceAlreadyExists
message: fmt.Sprintf("state of applicability with name %q already exists", s.Name),
}
} }
} }
return fmt.Errorf("cannot update state_of_applicability: %w", err) return fmt.Errorf("cannot update state_of_applicability: %w", err)
} }
if result.RowsAffected() == 0 { if result.RowsAffected() == 0 {
return &ErrStateOfApplicabilityNotFound{Identifier: s.ID.String()} return ErrResourceNotFound
} }
return nil return nil
@@ -325,7 +305,7 @@ WHERE %s
} }
if result.RowsAffected() == 0 { if result.RowsAffected() == 0 {
return &ErrStateOfApplicabilityNotFound{Identifier: s.ID.String()} return ErrResourceNotFound
} }
return nil return nil

View File

@@ -46,24 +46,8 @@ type (
} }
Tasks []*Task Tasks []*Task
ErrTaskNotFound struct {
Identifier string
}
ErrTaskAlreadyExists struct {
message string
}
) )
func (e ErrTaskNotFound) Error() string {
return fmt.Sprintf("task not found: %q", e.Identifier)
}
func (e ErrTaskAlreadyExists) Error() string {
return e.message
}
func (c Task) CursorKey(orderBy TaskOrderField) page.CursorKey { func (c Task) CursorKey(orderBy TaskOrderField) page.CursorKey {
switch orderBy { switch orderBy {
case TaskOrderFieldCreatedAt: case TaskOrderFieldCreatedAt:
@@ -114,7 +98,7 @@ LIMIT 1;
task, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[Task]) task, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[Task])
if err != nil { if err != nil {
if errors.Is(err, pgx.ErrNoRows) { if errors.Is(err, pgx.ErrNoRows) {
return &ErrTaskNotFound{Identifier: taskID.String()} return ErrResourceNotFound
} }
return fmt.Errorf("cannot collect tasks: %w", err) return fmt.Errorf("cannot collect tasks: %w", err)
@@ -185,9 +169,7 @@ VALUES (
var pgErr *pgconn.PgError var pgErr *pgconn.PgError
if errors.As(err, &pgErr) { if errors.As(err, &pgErr) {
if pgErr.Code == "23505" && pgErr.ConstraintName == "tasks_reference_id_unique" { if pgErr.Code == "23505" && pgErr.ConstraintName == "tasks_reference_id_unique" {
return &ErrTaskAlreadyExists{ return ErrResourceAlreadyExists
message: fmt.Sprintf("task with measure_id %s and reference_id %q already exists", c.MeasureID, c.ReferenceID),
}
} }
} }
return fmt.Errorf("cannot insert task: %w", err) return fmt.Errorf("cannot insert task: %w", err)

View File

@@ -21,11 +21,11 @@ import (
"maps" "maps"
"time" "time"
"go.probo.inc/probo/pkg/gid"
"go.probo.inc/probo/pkg/page"
"github.com/jackc/pgx/v5" "github.com/jackc/pgx/v5"
"github.com/jackc/pgx/v5/pgconn" "github.com/jackc/pgx/v5/pgconn"
"go.gearno.de/kit/pg" "go.gearno.de/kit/pg"
"go.probo.inc/probo/pkg/gid"
"go.probo.inc/probo/pkg/page"
) )
type ( type (
@@ -41,24 +41,8 @@ type (
} }
TrustCenters []*TrustCenter TrustCenters []*TrustCenter
ErrTrustCenterNotFound struct {
Identifier string
}
ErrTrustCenterAlreadyExists struct {
message string
}
) )
func (e ErrTrustCenterNotFound) Error() string {
return fmt.Sprintf("trust center not found: %q", e.Identifier)
}
func (e ErrTrustCenterAlreadyExists) Error() string {
return e.message
}
func (tc *TrustCenter) CursorKey(orderBy TrustCenterOrderField) page.CursorKey { func (tc *TrustCenter) CursorKey(orderBy TrustCenterOrderField) page.CursorKey {
switch orderBy { switch orderBy {
case TrustCenterOrderFieldCreatedAt: case TrustCenterOrderFieldCreatedAt:
@@ -239,9 +223,7 @@ INSERT INTO trust_centers (
var pgErr *pgconn.PgError var pgErr *pgconn.PgError
if errors.As(err, &pgErr) { if errors.As(err, &pgErr) {
if pgErr.Code == "23505" && pgErr.ConstraintName == "trust_centers_slug_key" { if pgErr.Code == "23505" && pgErr.ConstraintName == "trust_centers_slug_key" {
return &ErrTrustCenterAlreadyExists{ return ErrResourceAlreadyExists
message: fmt.Sprintf("trust center with slug %q already exists", tc.Slug),
}
} }
} }
return fmt.Errorf("cannot insert trust center: %w", err) return fmt.Errorf("cannot insert trust center: %w", err)

View File

@@ -48,24 +48,8 @@ type (
} }
TrustCenterAccesses []*TrustCenterAccess TrustCenterAccesses []*TrustCenterAccess
ErrTrustCenterAccessNotFound struct {
Identifier string
}
ErrTrustCenterAccessAlreadyExists struct {
message string
}
) )
func (e ErrTrustCenterAccessNotFound) Error() string {
return fmt.Sprintf("trust center access not found: %s", e.Identifier)
}
func (e ErrTrustCenterAccessAlreadyExists) Error() string {
return e.message
}
func (tca *TrustCenterAccess) CursorKey(orderBy TrustCenterAccessOrderField) page.CursorKey { func (tca *TrustCenterAccess) CursorKey(orderBy TrustCenterAccessOrderField) page.CursorKey {
switch orderBy { switch orderBy {
case TrustCenterAccessOrderFieldCreatedAt: case TrustCenterAccessOrderFieldCreatedAt:
@@ -117,7 +101,7 @@ LIMIT 1;
access, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[TrustCenterAccess]) access, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[TrustCenterAccess])
if err != nil { if err != nil {
if errors.Is(err, pgx.ErrNoRows) { if errors.Is(err, pgx.ErrNoRows) {
return &ErrTrustCenterAccessNotFound{Identifier: accessID.String()} return ErrResourceNotFound
} }
return fmt.Errorf("cannot collect trust center access: %w", err) return fmt.Errorf("cannot collect trust center access: %w", err)
@@ -175,7 +159,7 @@ LIMIT 1;
access, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[TrustCenterAccess]) access, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[TrustCenterAccess])
if err != nil { if err != nil {
if errors.Is(err, pgx.ErrNoRows) { if errors.Is(err, pgx.ErrNoRows) {
return &ErrTrustCenterAccessNotFound{Identifier: fmt.Sprintf("trust_center_id=%s, email=%s", trustCenterID, email)} return ErrResourceNotFound
} }
return fmt.Errorf("cannot collect trust center access: %w", err) return fmt.Errorf("cannot collect trust center access: %w", err)
@@ -235,9 +219,7 @@ INSERT INTO trust_center_accesses (
var pgErr *pgconn.PgError var pgErr *pgconn.PgError
if errors.As(err, &pgErr) { if errors.As(err, &pgErr) {
if pgErr.Code == "23505" && pgErr.ConstraintName == "trust_center_accesses_trust_center_id_email_key" { if pgErr.Code == "23505" && pgErr.ConstraintName == "trust_center_accesses_trust_center_id_email_key" {
return &ErrTrustCenterAccessAlreadyExists{ return ErrResourceAlreadyExists
message: "trust center access already exists",
}
} }
} }
return fmt.Errorf("cannot insert trust center access: %w", err) return fmt.Errorf("cannot insert trust center access: %w", err)

View File

@@ -42,24 +42,8 @@ type (
} }
TrustCenterDocumentAccesses []*TrustCenterDocumentAccess TrustCenterDocumentAccesses []*TrustCenterDocumentAccess
ErrTrustCenterDocumentAccessNotFound struct {
Identifier string
}
ErrTrustCenterDocumentAccessAlreadyExists struct {
message string
}
) )
func (e ErrTrustCenterDocumentAccessNotFound) Error() string {
return fmt.Sprintf("trust center document access not found: %s", e.Identifier)
}
func (e ErrTrustCenterDocumentAccessAlreadyExists) Error() string {
return e.message
}
func (tcda *TrustCenterDocumentAccess) CursorKey(orderBy TrustCenterDocumentAccessOrderField) page.CursorKey { func (tcda *TrustCenterDocumentAccess) CursorKey(orderBy TrustCenterDocumentAccessOrderField) page.CursorKey {
switch orderBy { switch orderBy {
case TrustCenterDocumentAccessOrderFieldCreatedAt: case TrustCenterDocumentAccessOrderFieldCreatedAt:
@@ -107,7 +91,7 @@ LIMIT 1;
access, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[TrustCenterDocumentAccess]) access, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[TrustCenterDocumentAccess])
if err != nil { if err != nil {
if errors.Is(err, pgx.ErrNoRows) { if errors.Is(err, pgx.ErrNoRows) {
return &ErrTrustCenterDocumentAccessNotFound{Identifier: accessID.String()} return ErrResourceNotFound
} }
return fmt.Errorf("cannot collect trust center document access: %w", err) return fmt.Errorf("cannot collect trust center document access: %w", err)
} }
@@ -267,18 +251,10 @@ INSERT INTO trust_center_document_accesses (
if errors.As(err, &pgErr) { if errors.As(err, &pgErr) {
if pgErr.Code == "23505" { if pgErr.Code == "23505" {
switch pgErr.ConstraintName { switch pgErr.ConstraintName {
case "trust_center_document_accesse_trust_center_access_id_docume_key": case "trust_center_document_accesse_trust_center_access_id_docume_key",
return &ErrTrustCenterDocumentAccessAlreadyExists{ "trust_center_document_accesse_trust_center_access_id_report_key",
message: fmt.Sprintf("trust center document access with trust_center_access_id %s and document_id %s already exists", tcda.TrustCenterAccessID, tcda.DocumentID), "trust_center_document_accesses_trust_center_file_id_key":
} return ErrResourceAlreadyExists
case "trust_center_document_accesse_trust_center_access_id_report_key":
return &ErrTrustCenterDocumentAccessAlreadyExists{
message: fmt.Sprintf("trust center document access with trust_center_access_id %s and report_id %s already exists", tcda.TrustCenterAccessID, tcda.ReportID),
}
case "trust_center_document_accesses_trust_center_file_id_key":
return &ErrTrustCenterDocumentAccessAlreadyExists{
message: fmt.Sprintf("trust center document access with trust_center_access_id %s and trust_center_file_id %s already exists", tcda.TrustCenterAccessID, tcda.TrustCenterFileID),
}
} }
} }
} }

View File

@@ -43,24 +43,8 @@ type (
} }
TrustCenterReferences []*TrustCenterReference TrustCenterReferences []*TrustCenterReference
ErrTrustCenterReferenceNotFound struct {
Identifier string
}
ErrTrustCenterReferenceAlreadyExists struct {
message string
}
) )
func (e ErrTrustCenterReferenceNotFound) Error() string {
return fmt.Sprintf("trust center reference not found: %q", e.Identifier)
}
func (e ErrTrustCenterReferenceAlreadyExists) Error() string {
return e.message
}
func (t TrustCenterReference) CursorKey(orderBy TrustCenterReferenceOrderField) page.CursorKey { func (t TrustCenterReference) CursorKey(orderBy TrustCenterReferenceOrderField) page.CursorKey {
switch orderBy { switch orderBy {
case TrustCenterReferenceOrderFieldRank: case TrustCenterReferenceOrderFieldRank:
@@ -174,9 +158,7 @@ RETURNING rank;
var pgErr *pgconn.PgError var pgErr *pgconn.PgError
if errors.As(err, &pgErr) { if errors.As(err, &pgErr) {
if pgErr.Code == "23505" && pgErr.ConstraintName == "trust_center_references_trust_center_id_rank_key" { if pgErr.Code == "23505" && pgErr.ConstraintName == "trust_center_references_trust_center_id_rank_key" {
return &ErrTrustCenterReferenceAlreadyExists{ return ErrResourceAlreadyExists
message: fmt.Sprintf("trust center reference with trust_center_id %s and rank already exists", t.TrustCenterID),
}
} }
} }
return fmt.Errorf("cannot insert trust center reference: %w", err) return fmt.Errorf("cannot insert trust center reference: %w", err)
@@ -221,7 +203,7 @@ WHERE
} }
if result.RowsAffected() == 0 { if result.RowsAffected() == 0 {
return ErrTrustCenterReferenceNotFound{Identifier: t.ID.String()} return ErrResourceNotFound
} }
return nil return nil

View File

@@ -43,24 +43,8 @@ type (
} }
Users []*User Users []*User
ErrUserNotFound struct {
Identifier string
}
ErrUserAlreadyExists struct {
message string
}
) )
func (e ErrUserNotFound) Error() string {
return fmt.Sprintf("user not found: %q", e.Identifier)
}
func (e ErrUserAlreadyExists) Error() string {
return e.message
}
func (u User) CursorKey(orderBy UserOrderField) page.CursorKey { func (u User) CursorKey(orderBy UserOrderField) page.CursorKey {
switch orderBy { switch orderBy {
case UserOrderFieldCreatedAt: case UserOrderFieldCreatedAt:
@@ -180,7 +164,7 @@ LIMIT 1;
user, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[User]) user, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[User])
if err != nil { if err != nil {
if errors.Is(err, pgx.ErrNoRows) { if errors.Is(err, pgx.ErrNoRows) {
return &ErrUserNotFound{Identifier: email.String()} return ErrResourceNotFound
} }
return fmt.Errorf("cannot collect user: %w", err) return fmt.Errorf("cannot collect user: %w", err)
@@ -224,7 +208,7 @@ LIMIT 1;
user, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[User]) user, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[User])
if err != nil { if err != nil {
if errors.Is(err, pgx.ErrNoRows) { if errors.Is(err, pgx.ErrNoRows) {
return &ErrUserNotFound{Identifier: userID.String()} return ErrResourceNotFound
} }
return fmt.Errorf("cannot collect user: %w", err) return fmt.Errorf("cannot collect user: %w", err)
@@ -238,7 +222,6 @@ LIMIT 1;
func (u *User) Insert( func (u *User) Insert(
ctx context.Context, ctx context.Context,
conn pg.Conn, conn pg.Conn,
scope Scoper,
) error { ) error {
q := ` q := `
INSERT INTO INSERT INTO
@@ -272,9 +255,7 @@ VALUES (
var pgErr *pgconn.PgError var pgErr *pgconn.PgError
if errors.As(err, &pgErr) { if errors.As(err, &pgErr) {
if pgErr.Code == "23505" && strings.Contains(pgErr.ConstraintName, "email_address") { if pgErr.Code == "23505" && strings.Contains(pgErr.ConstraintName, "email_address") {
return &ErrUserAlreadyExists{ return ErrResourceAlreadyExists
message: fmt.Sprintf("user with email %s already exists", u.EmailAddress),
}
} }
} }
@@ -284,71 +265,6 @@ VALUES (
return nil return nil
} }
func (u *User) UpdateEmailVerification(
ctx context.Context,
conn pg.Conn,
verified bool,
) error {
q := `
UPDATE
users
SET
email_address_verified = @email_address_verified,
updated_at = @updated_at
WHERE
id = @user_id
`
args := pgx.StrictNamedArgs{
"user_id": u.ID,
"email_address_verified": verified,
"updated_at": time.Now(),
}
_, err := conn.Exec(ctx, q, args)
if err != nil {
return fmt.Errorf("cannot update user email verification: %w", err)
}
u.EmailAddressVerified = verified
u.UpdatedAt = args["updated_at"].(time.Time)
return nil
}
func (u *User) UpdatePassword(
ctx context.Context,
conn pg.Conn,
hashedPassword []byte,
) error {
q := `
UPDATE
users
SET
hashed_password = @hashed_password,
updated_at = @updated_at
WHERE
id = @user_id
`
now := time.Now()
args := pgx.StrictNamedArgs{
"user_id": u.ID,
"hashed_password": hashedPassword,
"updated_at": now,
}
_, err := conn.Exec(ctx, q, args)
if err != nil {
return fmt.Errorf("cannot update user password: %w", err)
}
u.HashedPassword = hashedPassword
u.UpdatedAt = now
return nil
}
func (u *User) Update(ctx context.Context, conn pg.Conn) error { func (u *User) Update(ctx context.Context, conn pg.Conn) error {
q := ` q := `
UPDATE UPDATE
@@ -357,6 +273,8 @@ SET
email_address = @email_address, email_address = @email_address,
email_address_verified = @email_address_verified, email_address_verified = @email_address_verified,
saml_subject = @saml_subject, saml_subject = @saml_subject,
fullname = @fullname,
hashed_password = @hashed_password,
updated_at = @updated_at updated_at = @updated_at
WHERE WHERE
id = @user_id id = @user_id
@@ -368,13 +286,19 @@ WHERE
"email_address_verified": u.EmailAddressVerified, "email_address_verified": u.EmailAddressVerified,
"saml_subject": u.SAMLSubject, "saml_subject": u.SAMLSubject,
"updated_at": u.UpdatedAt, "updated_at": u.UpdatedAt,
"fullname": u.FullName,
"hashed_password": u.HashedPassword,
} }
_, err := conn.Exec(ctx, q, args) result, err := conn.Exec(ctx, q, args)
if err != nil { if err != nil {
return fmt.Errorf("cannot update user: %w", err) return fmt.Errorf("cannot update user: %w", err)
} }
if result.RowsAffected() == 0 {
return ErrResourceNotFound
}
return nil return nil
} }
@@ -411,7 +335,7 @@ LIMIT 1;
user, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[User]) user, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[User])
if err != nil { if err != nil {
if errors.Is(err, pgx.ErrNoRows) { if errors.Is(err, pgx.ErrNoRows) {
return &ErrUserNotFound{Identifier: samlSubject} return ErrResourceNotFound
} }
return fmt.Errorf("cannot collect user: %w", err) return fmt.Errorf("cannot collect user: %w", err)
@@ -422,10 +346,6 @@ LIMIT 1;
return nil return nil
} }
// LoadByEmailAndTenant, LoadByEmailGlobal, and IsTenantUser methods removed
// All users are now global (no tenant_id distinction)
// Use LoadByEmail() for all email-based lookups
func (u *User) CountMemberships( func (u *User) CountMemberships(
ctx context.Context, ctx context.Context,
conn pg.Conn, conn pg.Conn,

View File

@@ -23,27 +23,30 @@ import (
"github.com/jackc/pgx/v5" "github.com/jackc/pgx/v5"
"go.gearno.de/kit/pg" "go.gearno.de/kit/pg"
"go.probo.inc/probo/pkg/gid" "go.probo.inc/probo/pkg/gid"
"go.probo.inc/probo/pkg/page"
) )
type ( type (
UserAPIKey struct { UserAPIKey struct {
ID gid.GID `db:"id"` ID gid.GID `db:"id"`
UserID gid.GID `db:"user_id"` UserID gid.GID `db:"user_id"`
Name string `db:"name"` Name string `db:"name"`
ExpiresAt time.Time `db:"expires_at"` ExpiresAt time.Time `db:"expires_at"`
CreatedAt time.Time `db:"created_at"` ExpireReason *ExpireReason `db:"expire_reason"`
UpdatedAt time.Time `db:"updated_at"` CreatedAt time.Time `db:"created_at"`
UpdatedAt time.Time `db:"updated_at"`
} }
UserAPIKeys []*UserAPIKey UserAPIKeys []*UserAPIKey
ErrUserAPIKeyNotFound struct {
Identifier string
}
) )
func (e ErrUserAPIKeyNotFound) Error() string { func (a *UserAPIKey) CursorKey(orderBy UserAPIKeyOrderField) page.CursorKey {
return fmt.Sprintf("user api key not found: %q", e.Identifier) switch orderBy {
case UserAPIKeyOrderFieldCreatedAt:
return page.NewCursorKey(a.ID, a.CreatedAt)
}
panic(fmt.Sprintf("unsupported order by: %s", orderBy))
} }
func (a *UserAPIKey) LoadByID( func (a *UserAPIKey) LoadByID(
@@ -57,6 +60,7 @@ SELECT
user_id, user_id,
name, name,
expires_at, expires_at,
expire_reason,
created_at, created_at,
updated_at updated_at
FROM FROM
@@ -76,7 +80,7 @@ LIMIT 1;
apiKey, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[UserAPIKey]) apiKey, err := pgx.CollectExactlyOneRow(rows, pgx.RowToStructByName[UserAPIKey])
if err != nil { if err != nil {
if errors.Is(err, pgx.ErrNoRows) { if errors.Is(err, pgx.ErrNoRows) {
return &ErrUserAPIKeyNotFound{Identifier: apiKeyID.String()} return ErrResourceNotFound
} }
return fmt.Errorf("cannot collect user api key: %w", err) return fmt.Errorf("cannot collect user api key: %w", err)
@@ -98,6 +102,7 @@ SELECT
user_id, user_id,
name, name,
expires_at, expires_at,
expire_reason,
created_at, created_at,
updated_at updated_at
FROM FROM
@@ -124,30 +129,53 @@ ORDER BY created_at DESC;
return nil return nil
} }
func (a *UserAPIKeys) CountByUserID(ctx context.Context, conn pg.Conn, userID gid.GID) (int, error) {
q := `
SELECT
COUNT(*)
FROM
auth_user_api_keys
WHERE
user_id = @user_id
ORDER BY created_at DESC;
`
args := pgx.StrictNamedArgs{"user_id": userID}
row := conn.QueryRow(ctx, q, args)
var count int
if err := row.Scan(&count); err != nil {
return 0, fmt.Errorf("cannot scan count: %w", err)
}
return count, nil
}
func (a *UserAPIKey) Insert( func (a *UserAPIKey) Insert(
ctx context.Context, ctx context.Context,
conn pg.Conn, conn pg.Conn,
) error { ) error {
q := ` q := `
INSERT INTO INSERT INTO
auth_user_api_keys (id, user_id, name, expires_at, created_at, updated_at) auth_user_api_keys (id, user_id, name, expires_at, expire_reason, created_at, updated_at)
VALUES ( VALUES (
@api_key_id, @api_key_id,
@user_id, @user_id,
@name, @name,
@expires_at, @expires_at,
@expire_reason,
@created_at, @created_at,
@updated_at @updated_at
) )
` `
args := pgx.StrictNamedArgs{ args := pgx.StrictNamedArgs{
"api_key_id": a.ID, "api_key_id": a.ID,
"user_id": a.UserID, "user_id": a.UserID,
"name": a.Name, "name": a.Name,
"expires_at": a.ExpiresAt, "expires_at": a.ExpiresAt,
"created_at": a.CreatedAt, "expire_reason": a.ExpireReason,
"updated_at": a.UpdatedAt, "created_at": a.CreatedAt,
"updated_at": a.UpdatedAt,
} }
_, err := conn.Exec(ctx, q, args) _, err := conn.Exec(ctx, q, args)
@@ -168,16 +196,18 @@ UPDATE
SET SET
name = @name, name = @name,
expires_at = @expires_at, expires_at = @expires_at,
expire_reason = @expire_reason,
updated_at = @updated_at updated_at = @updated_at
WHERE WHERE
id = @api_key_id id = @api_key_id
` `
args := pgx.StrictNamedArgs{ args := pgx.StrictNamedArgs{
"api_key_id": a.ID, "api_key_id": a.ID,
"name": a.Name, "name": a.Name,
"expires_at": a.ExpiresAt, "expires_at": a.ExpiresAt,
"updated_at": a.UpdatedAt, "expire_reason": a.ExpireReason,
"updated_at": a.UpdatedAt,
} }
_, err := conn.Exec(ctx, q, args) _, err := conn.Exec(ctx, q, args)

Some files were not shown because too many files have changed in this diff Show More