update (most tests pass)

This commit is contained in:
Evgeny Poberezkin
2024-11-10 12:30:24 +00:00
parent 28105038d4
commit 90ed503ee0
6 changed files with 122 additions and 121 deletions
+26 -28
View File
@@ -106,7 +106,7 @@ import qualified Simplex.Messaging.Crypto as C
import qualified Simplex.Messaging.Crypto.Ratchet as CR
import Simplex.Messaging.Encoding.String
import Simplex.Messaging.Parsers (defaultJSON)
import Simplex.Messaging.Protocol (BasicAuth (..), ProtoServerWithAuth (..), ProtocolServer (..), ProtocolTypeI (..), SubscriptionMode)
import Simplex.Messaging.Protocol (BasicAuth (..), ProtoServerWithAuth (..), ProtocolServer (..), ProtocolTypeI (..), SProtocolType (..), SubscriptionMode, UserProtocol)
import Simplex.Messaging.Transport.Client (TransportHost)
import Simplex.Messaging.Util (eitherToMaybe, safeDecodeUtf8)
@@ -530,19 +530,19 @@ updateUserAddressAutoAccept db user@User {userId} autoAccept = do
Just AutoAccept {acceptIncognito, autoReply} -> (True, acceptIncognito, autoReply)
_ -> (False, False, Nothing)
getUpdateUserServers :: forall p. ProtocolTypeI p => DB.Connection -> NonEmpty PresetOperator -> NonEmpty (NewUserServer p) -> User -> IO (NonEmpty (UserServer p))
getUpdateUserServers db presetOps randomSrvs user = do
getUpdateUserServers :: forall p. (ProtocolTypeI p, UserProtocol p) => DB.Connection -> SProtocolType p -> NonEmpty PresetOperator -> NonEmpty (NewUserServer p) -> User -> IO (NonEmpty (UserServer p))
getUpdateUserServers db p presetOps randomSrvs user = do
ts <- getCurrentTime
srvs <- getProtocolServers db user
let srvs' = updatedUserServers presetOps randomSrvs srvs
let srvs' = updatedUserServers p presetOps randomSrvs srvs
mapM (upsertServer ts) srvs'
where
upsertServer :: UTCTime -> AUserServer p -> IO (UserServer p)
upsertServer ts (AUS _ s@UserServer {serverId}) = case serverId of
DBNewEntity -> insertProtocolServer db user ts s
DBEntityId _ -> updateServer s ts $> s
updateServer :: UserServer p -> UTCTime -> IO ()
updateServer UserServer {serverId, server, preset, tested, enabled} ts =
DBNewEntity -> insertProtocolServer db p user ts s
DBEntityId _ -> updateServer ts s $> s
updateServer :: UTCTime -> UserServer p -> IO ()
updateServer ts UserServer {serverId, server, preset, tested, enabled} =
DB.execute
db
[sql|
@@ -551,7 +551,7 @@ getUpdateUserServers db presetOps randomSrvs user = do
preset = ?, tested = ?, enabled = ?, updated_at
WHERE smp_server_id = ?
|]
(serverColumns server :. (preset, tested, enabled, ts, serverId))
(serverColumns p server :. (preset, tested, enabled, ts, serverId))
getProtocolServers :: forall p. ProtocolTypeI p => DB.Connection -> User -> IO [UserServer p]
getProtocolServers db User {userId} =
@@ -574,12 +574,12 @@ getProtocolServers db User {userId} =
-- TODO remove
-- overwriteOperatorsAndServers :: forall p. ProtocolTypeI p => DB.Connection -> User -> Maybe [ServerOperator] -> [ServerCfg p] -> ExceptT StoreError IO [ServerCfg p]
-- overwriteOperatorsAndServers db user@User {userId} operators_ servers = do
overwriteProtocolServers :: forall p. ProtocolTypeI p => DB.Connection -> User -> [UserServer p] -> ExceptT StoreError IO ()
overwriteProtocolServers db User {userId} servers =
overwriteProtocolServers :: ProtocolTypeI p => DB.Connection -> SProtocolType p -> User -> [UserServer p] -> ExceptT StoreError IO ()
overwriteProtocolServers db p User {userId} servers =
-- liftIO $ mapM_ (updateServerOperators_ db) operators_
checkConstraint SEUniqueID . ExceptT $ do
currentTs <- getCurrentTime
DB.execute db "DELETE FROM protocol_servers WHERE user_id = ? AND protocol = ? " (userId, protocol)
DB.execute db "DELETE FROM protocol_servers WHERE user_id = ? AND protocol = ? " (userId, decodeLatin1 $ strEncode p)
forM_ servers $ \UserServer {serverId, server, preset, tested, enabled} -> do
DB.execute
db
@@ -588,13 +588,11 @@ overwriteProtocolServers db User {userId} servers =
(server_id, protocol, host, port, key_hash, basic_auth, preset, tested, enabled, user_id, created_at, updated_at)
VALUES (?,?,?,?,?,?,?,?,?,?,?,?)
|]
(Only serverId :. serverColumns server :. (preset, tested, enabled, userId, currentTs, currentTs))
(Only serverId :. serverColumns p server :. (preset, tested, enabled, userId, currentTs, currentTs))
pure $ Right ()
where
protocol = decodeLatin1 $ strEncode $ protocolTypeI @p
insertProtocolServer :: forall p. ProtocolTypeI p => DB.Connection -> User -> UTCTime -> NewUserServer p -> IO (UserServer p)
insertProtocolServer db User {userId} ts srv@UserServer {server, preset, tested, enabled} = do
insertProtocolServer :: forall p. ProtocolTypeI p => DB.Connection -> SProtocolType p -> User -> UTCTime -> NewUserServer p -> IO (UserServer p)
insertProtocolServer db p User {userId} ts srv@UserServer {server, preset, tested, enabled} = do
DB.execute
db
[sql|
@@ -602,13 +600,13 @@ insertProtocolServer db User {userId} ts srv@UserServer {server, preset, tested,
(protocol, host, port, key_hash, basic_auth, preset, tested, enabled, user_id, created_at, updated_at)
VALUES (?,?,?,?,?,?,?,?,?,?,?)
|]
(serverColumns server :. (preset, tested, enabled, userId, ts, ts))
(serverColumns p server :. (preset, tested, enabled, userId, ts, ts))
sId <- insertedRowId db
pure (srv :: NewUserServer p) {serverId = DBEntityId sId}
serverColumns :: forall p. ProtocolTypeI p => ProtoServerWithAuth p -> (Text, NonEmpty TransportHost, String, C.KeyHash, Maybe Text)
serverColumns (ProtoServerWithAuth ProtocolServer {host, port, keyHash} auth_) =
let protocol = decodeLatin1 $ strEncode $ protocolTypeI @p
serverColumns :: ProtocolTypeI p => SProtocolType p -> ProtoServerWithAuth p -> (Text, NonEmpty TransportHost, String, C.KeyHash, Maybe Text)
serverColumns p (ProtoServerWithAuth ProtocolServer {host, port, keyHash} auth_) =
let protocol = decodeLatin1 $ strEncode p
auth = safeDecodeUtf8 . unBasicAuth <$> auth_
in (protocol, host, port, keyHash, auth)
@@ -694,7 +692,7 @@ getServerOperators_ db =
db
[sql|
SELECT server_operator_id, server_operator_tag, trade_name, legal_name,
server_domains, enabled, role_storage, role_proxy,
server_domains, enabled, role_storage, role_proxy
FROM server_operators
|]
where
@@ -807,8 +805,8 @@ setUserServers db User {userId} userServers = do
forM_ userServers $ do
\UserOperatorServers {operator, smpServers, xftpServers} -> do
forM_ operator $ \op -> liftIO $ updateOperator currentTs op
overwriteServers currentTs operator smpServers
overwriteServers currentTs operator xftpServers
overwriteServers SPSMP currentTs operator smpServers
overwriteServers SPXFTP currentTs operator xftpServers
where
updateOperator :: UTCTime -> ServerOperator -> IO ()
updateOperator currentTs ServerOperator {operatorId, enabled, roles = ServerRoles {storage, proxy}} =
@@ -820,8 +818,8 @@ setUserServers db User {userId} userServers = do
WHERE server_operator_id = ?
|]
(enabled, storage, proxy, operatorId, currentTs)
overwriteServers :: forall p. ProtocolTypeI p => UTCTime -> Maybe ServerOperator -> [UserServer p] -> ExceptT StoreError IO ()
overwriteServers currentTs serverOperator servers =
overwriteServers :: ProtocolTypeI p => SProtocolType p -> UTCTime -> Maybe ServerOperator -> [UserServer p] -> ExceptT StoreError IO ()
overwriteServers p currentTs serverOperator servers =
checkConstraint SEUniqueID . ExceptT $ do
case serverOperator of
Nothing ->
@@ -836,11 +834,11 @@ setUserServers db User {userId} userServers = do
(server_id, protocol, host, port, key_hash, basic_auth, preset, tested, enabled, user_id, created_at, updated_at)
VALUES (?,?,?,?,?,?,?,?,?,?,?,?,?)
|]
(Only serverId :. serverColumns server :. (tested, enabled, userId, currentTs, currentTs))
(Only serverId :. serverColumns p server :. (tested, enabled, userId, currentTs, currentTs))
-- take preset from operator
pure $ Right ()
where
protocol = decodeLatin1 $ strEncode $ protocolTypeI @p
protocol = decodeLatin1 $ strEncode p
createCall :: DB.Connection -> User -> Call -> UTCTime -> IO ()
createCall db user@User {userId} Call {contactId, callId, callUUID, chatItemId, callState} callTs = do